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..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
@@ -34,12 +34,17 @@
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.nio.file.Paths;
+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 +62,168 @@ private MtlsUtils() {
// Prevent instantiation for Utility class
}
+ /**
+ * 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) {
+ 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:
+ *
+ *
+ *
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.
+ *
{@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.
+ *
{@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.
+ *
If the certificate on disk differs from the active one, this refresh is treated like a
+ * rotation refresh: channels that fail to be recreated are dropped (and the pool refilled), so
+ * the switch is completed and the generation advanced only once every channel left in the pool
+ * uses the new certificate. Otherwise a channel that fails to be recreated keeps its slot, and
+ * the generation is not advanced, 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 {
- refresh();
+ // See refresh(): a generation change while waiting for the lock means a concurrent refresh
+ // just recreated every channel on a new certificate.
+ long generationBeforeLock = generation.get();
+ synchronized (entryWriteLock) {
+ if (workloadCertPath == null) {
+ refreshAll();
+ return;
+ }
+ if (generation.get() != generationBeforeLock) {
+ LOG.fine(
+ "Skipping pre-emptive channel refresh: channels were just recreated by a concurrent"
+ + " refresh");
+ return;
+ }
+ String currentDiskFingerprint = rotationTracker.readDiskFingerprint();
+ if (currentDiskFingerprint.isEmpty()) {
+ // The configured certificate could not be read, which normally means it is being
+ // rewritten during a rotation. Skip this refresh rather than recreate channels from a
+ // partially written certificate; the next periodic or reactive refresh recreates them.
+ LOG.fine(
+ "Skipping pre-emptive channel refresh: the workload certificate could not be read;"
+ + " channels will be recreated on the next refresh");
+ return;
+ }
+ boolean certChanged = !rotationTracker.isAlreadyActive(currentDiskFingerprint);
+ if (refreshAll(/* dropUnrefreshedChannels= */ certChanged) && certChanged) {
+ completeCertificateSwitch(currentDiskFingerprint);
+ }
+ }
} 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);
}
}
+ @VisibleForTesting
+ void invalidateDiskFingerprintCache() {
+ rotationTracker.invalidateCache();
+ }
+
+ boolean shouldRefresh() {
+ return rotationTracker.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() {
+ // A generation change while waiting for the lock means a concurrent refresh already switched
+ // every channel to a new certificate, so there is no need to read the certificate again.
+ long generationBeforeLock = generation.get();
// Note: synchronization is necessary in case refresh is called concurrently:
// - thread1 fails to replace a single entry
// - thread2 succeeds replacing an entry
@@ -443,28 +532,179 @@ 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;
+ }
+ if (generation.get() != generationBeforeLock) {
+ LOG.fine(
+ "Channel pool was already refreshed by a concurrent thread, skipping duplicate"
+ + " refresh");
+ return;
+ }
+ String currentDiskFingerprint = rotationTracker.readDiskFingerprint();
+ if (currentDiskFingerprint.isEmpty()) {
+ return;
+ }
+
+ // Double-check fingerprint inside the lock
+ if (rotationTracker.isAlreadyActive(currentDiskFingerprint)) {
+ LOG.fine(
+ "Channel pool was already refreshed by a concurrent thread, skipping duplicate"
+ + " refresh");
+ return;
+ }
+
+ // Drop any channel that fails to refresh so that no traffic is routed to the old certificate.
+ if (refreshAll(/* dropUnrefreshedChannels= */ true)) {
+ 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;
+ }
+
+ @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, and the
+ * pool is then refilled asynchronously to its size before the refresh.
+ * @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;
+ }
LOG.fine("Refreshing all channels");
- ArrayList newEntries = new ArrayList<>(entries.get());
+ // All writes to entries happen under entryWriteLock, so this snapshot is the current pool.
+ ImmutableList currentEntries = entries.get();
+ List keptEntries = new ArrayList<>(currentEntries.size());
+ List retiredEntries = new ArrayList<>(currentEntries.size());
+ List createdEntries = new ArrayList<>(currentEntries.size());
+ boolean allCreated = !currentEntries.isEmpty();
- for (int i = 0; i < newEntries.size(); i++) {
- try {
- newEntries.set(i, new Entry(channelFactory.createSingleChannel()));
- } catch (IOException e) {
- LOG.log(Level.WARNING, "Failed to refresh channel, leaving old channel", e);
+ try {
+ for (Entry oldEntry : currentEntries) {
+ try {
+ Entry newEntry = new Entry(channelFactory.createSingleChannel());
+ createdEntries.add(newEntry);
+ keptEntries.add(newEntry);
+ retiredEntries.add(oldEntry);
+ } catch (Exception e) {
+ allCreated = false;
+ if (dropUnrefreshedChannels) {
+ retiredEntries.add(oldEntry);
+ } else {
+ keptEntries.add(oldEntry);
+ }
+ LOG.log(
+ Level.WARNING,
+ dropUnrefreshedChannels
+ ? "Failed to refresh channel, dropping old channel"
+ : "Failed to refresh channel, leaving old channel",
+ e);
+ }
}
- }
- ImmutableList replacedEntries = entries.getAndSet(ImmutableList.copyOf(newEntries));
+ if (createdEntries.isEmpty()) {
+ return false;
+ }
+
+ entries.set(ImmutableList.copyOf(keptEntries));
+ createdEntries.clear(); // Ownership transferred to pool
+
+ // Shutdown the channels that were cycled out.
+ retiredEntries.forEach(Entry::requestShutdown);
- // Shutdown the channels that were cycled out.
- for (Entry e : replacedEntries) {
- if (!newEntries.contains(e)) {
+ // Restore the pool to its pre-refresh size right away, rather than leave a dynamically
+ // sized
+ // pool to grow back over several resize() runs while traffic queues on fewer channels.
+ if (dropUnrefreshedChannels && !allCreated) {
+ scheduleRefill(currentEntries.size());
+ }
+ return dropUnrefreshedChannels || allCreated;
+ } finally {
+ // If an Error aborted before the swap, shut down newly created channels so they don't leak
+ for (Entry e : createdEntries) {
e.requestShutdown();
}
}
}
}
+ /**
+ * Schedules a one-shot task that restores the pool to {@code targetSize} channels after a
+ * certificate rotation refresh dropped channels that failed to refresh.
+ */
+ private void scheduleRefill(int targetSize) {
+ try {
+ backgroundExecutorProvider.getExecutor().execute(() -> refillSafely(targetSize));
+ } catch (RuntimeException e) {
+ LOG.log(Level.WARNING, "Failed to schedule channel pool refill", e);
+ }
+ }
+
+ private void refillSafely(int targetSize) {
+ try {
+ synchronized (entryWriteLock) {
+ if (isShutdown) {
+ return;
+ }
+ 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.
+ *
+ *
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();
+ }
+
/**
* Get and retain a Channel Entry. The returned Entry will have its rpc count incremented,
* preventing it from getting recycled.
@@ -616,17 +856,30 @@ 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;
+ }
}
}
- /** 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)} 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;
+ 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();
+ private final AtomicBoolean wasStarted = new AtomicBoolean();
public ReleasingClientCall(ClientCall delegate, Entry entry) {
super(delegate);
@@ -635,51 +888,81 @@ public ReleasingClientCall(ClientCall delegate, Entry entry) {
@Override
public void start(Listener responseListener, Metadata headers) {
- if (cancellationException != null) {
- 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 {
+ 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 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);
- super.cancel(message, 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();
+ }
+ }
}
}
}
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..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
@@ -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. */
@@ -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,12 @@ public ApiCallContext merge(ApiCallContext inputCallContext) {
newCallOptions = newCallOptions.withOption(TRACER_KEY, newTracer);
}
+ TransportChannel newTransportChannel = grpcCallContext.transportChannel;
+ if (newTransportChannel == null
+ && (grpcCallContext.channel == null || grpcCallContext.channel.equals(channel))) {
+ newTransportChannel = transportChannel;
+ }
+
// The EndpointContext is not updated as there should be no reason for a user
// to update this.
return new GrpcCallContext(
@@ -557,7 +583,8 @@ public ApiCallContext merge(ApiCallContext inputCallContext) {
newRetrySettings,
newRetryableCodes,
endpointContext,
- newIsDirectPath);
+ newIsDirectPath,
+ newTransportChannel);
}
/** The {@link Channel} set on this context. */
@@ -635,7 +662,8 @@ public GrpcCallContext withChannel(@Nullable Channel newChannel) {
retrySettings,
retryableCodes,
endpointContext,
- isDirectPath);
+ isDirectPath,
+ (newChannel != null && newChannel.equals(channel)) ? transportChannel : null);
}
/** Returns a new instance with the call options set to the given call options. */
@@ -653,7 +681,8 @@ public GrpcCallContext withCallOptions(CallOptions newCallOptions) {
retrySettings,
retryableCodes,
endpointContext,
- isDirectPath);
+ isDirectPath,
+ transportChannel);
}
public GrpcCallContext withRequestParamsDynamicHeaderOption(String requestParams) {
@@ -698,7 +727,8 @@ public GrpcCallContext withOption(Key key, T value) {
retrySettings,
retryableCodes,
endpointContext,
- isDirectPath);
+ isDirectPath,
+ transportChannel);
}
/** {@inheritDoc} */
@@ -759,7 +789,8 @@ public int hashCode() {
options,
retrySettings,
retryableCodes,
- endpointContext);
+ endpointContext,
+ transportChannel);
}
@Override
@@ -783,7 +814,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..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
@@ -68,6 +68,32 @@ 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 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/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..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,12 +401,21 @@ 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.getWorkloadCertPath()
+ : null;
return GrpcTransportChannel.newBuilder()
.setManagedChannel(
ChannelPool.create(
channelPoolSettings,
InstantiatingGrpcChannelProvider.this::createSingleChannel,
- backgroundExecutor))
+ backgroundExecutor,
+ workloadCertPath))
.setDirectPath(this.canUseDirectPath())
.build();
}
@@ -465,8 +474,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 +484,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 +678,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 +698,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;
@@ -755,6 +770,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
@@ -1403,7 +1420,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..ec2bacff40e6 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,12 +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;
@@ -81,13 +86,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 +111,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 +128,7 @@ void testRoundRobin() throws IOException {
ChannelPool.create(
ChannelPoolSettings.staticallySized(channels.size()),
new FakeChannelFactory(channels),
+ null,
null);
verifyTargetChannel(pool, channels, sub1);
@@ -195,6 +207,7 @@ void ensureEvenDistribution() throws InterruptedException, IOException {
ChannelPool.create(
ChannelPoolSettings.staticallySized(numChannels),
new FakeChannelFactory(Arrays.asList(channels)),
+ null,
null);
int numThreads = 20;
@@ -233,6 +246,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 +287,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 +312,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 +337,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 +361,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 +388,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 +406,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 +437,978 @@ 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 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 channelReactiveMTlsRefresh_swapsChannelsOnlyWhenCertChanges()
+ 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
+ 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);
+
+ // 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 channelReactiveMTlsRefresh_failedCreationDoesNotMutateFingerprintAndAllowsRetry()
+ throws IOException {
+ ManagedChannel channel1 = Mockito.mock(ManagedChannel.class);
+ ManagedChannel channel2 = Mockito.mock(ManagedChannel.class);
+ ChannelFactory channelFactory =
+ Mockito.mock(ChannelFactory.class, Mockito.withSettings().withoutAnnotations());
+
+ // 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));
+ }
+
+ 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_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 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(refilled);
+
+ createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor);
+ long genBefore = pool.getGeneration();
+
+ pool.refresh();
+
+ // 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();
+ assertThat(pool.shouldRefresh()).isFalse();
+
+ // 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);
+
+ pool.refresh();
+
+ 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();
+ List before = pool.entries.get();
+
+ assertThat(pool.refreshAll()).isFalse();
+
+ assertThat(pool.entries.get()).hasSize(2);
+ // Each channel keeps its slot, so affinity-to-index mapping is unchanged.
+ assertThat(pool.entries.get().get(0)).isNotSameInstanceAs(before.get(0));
+ assertThat(pool.entries.get().get(1)).isSameInstanceAs(before.get(1));
+ Mockito.verify(initial1).shutdown();
+ Mockito.verify(initial2, Mockito.never()).shutdown();
+ assertThat(pool.getGeneration()).isEqualTo(genBefore);
+ Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class));
+ }
+
+ @Test
+ void channelReactiveMTlsRefresh_allChannelsFail_leavesPoolUnchangedAndDoesNotRefill()
+ throws IOException {
+ ScheduledExecutorService executor = mockExecutor();
+ ManagedChannel initial1 = Mockito.mock(ManagedChannel.class);
+ ManagedChannel initial2 = Mockito.mock(ManagedChannel.class);
+ ChannelFactory channelFactory = mockChannelFactory();
+ Mockito.when(channelFactory.createSingleChannel())
+ .thenReturn(initial1, initial2)
+ .thenThrow(new IOException("Failure on first sub-channel"))
+ .thenThrow(new IOException("Failure on second sub-channel"));
+
+ createMtlsPoolAndRotateCert(ChannelPoolSettings.staticallySized(2), channelFactory, executor);
+ long genBefore = pool.getGeneration();
+ List before = pool.entries.get();
+
+ pool.refresh();
+
+ // Nothing could be rebuilt, so the pool keeps its old channels and the rotation stays pending.
+ assertThat(pool.entries.get()).containsExactlyElementsIn(before).inOrder();
+ Mockito.verify(initial1, Mockito.never()).shutdown();
+ Mockito.verify(initial2, Mockito.never()).shutdown();
+ assertThat(pool.getGeneration()).isEqualTo(genBefore);
+ pool.invalidateDiskFingerprintCache();
+ assertThat(pool.shouldRefresh()).isTrue();
+ Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class));
+ }
+
+ @Test
+ void channelReactiveMTlsRefresh_partialFailureInDynamicPool_refillsToPreRefreshSize()
+ throws IOException {
+ ScheduledExecutorService executor = mockExecutor();
+ ManagedChannel initial1 = Mockito.mock(ManagedChannel.class);
+ ManagedChannel initial2 = Mockito.mock(ManagedChannel.class);
+ ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class);
+ 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(refilled);
+
+ createMtlsPoolAndRotateCert(
+ ChannelPoolSettings.builder()
+ .setInitialChannelCount(2)
+ .setMinChannelCount(2)
+ .setMaxChannelCount(4)
+ .setMinRpcsPerChannel(1)
+ .setMaxRpcsPerChannel(2)
+ .build(),
+ channelFactory,
+ executor);
+
+ pool.refresh();
+
+ assertThat(pool.entries.get()).hasSize(1);
+ Mockito.verify(initial1).shutdown();
+ Mockito.verify(initial2).shutdown();
+ pool.invalidateDiskFingerprintCache();
+ assertThat(pool.shouldRefresh()).isFalse();
+
+ // Dynamic pools are refilled right away too, rather than waiting for resize().
+ ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class);
+ Mockito.verify(executor).execute(refillTask.capture());
+ refillTask.getValue().run();
+ assertThat(pool.entries.get()).hasSize(2);
+ Mockito.verify(refilled, Mockito.never()).shutdown();
+ }
+
+ @Test
+ void channelReactiveMTlsRefresh_partialFailureAfterResize_refillsToPreRefreshSize()
+ throws IOException {
+ ScheduledExecutorService executor = mockExecutor();
+ ChannelFactory channelFactory = mockChannelFactory();
+ Mockito.when(channelFactory.createSingleChannel())
+ .thenReturn(
+ // initial pool of 4
+ Mockito.mock(ManagedChannel.class),
+ Mockito.mock(ManagedChannel.class),
+ Mockito.mock(ManagedChannel.class),
+ Mockito.mock(ManagedChannel.class),
+ // rotation refresh: 1 of the 2 remaining channels is recreated
+ Mockito.mock(ManagedChannel.class))
+ .thenThrow(new IOException("Transient failure on second sub-channel"))
+ .thenReturn(Mockito.mock(ManagedChannel.class));
+
+ createMtlsPoolAndRotateCert(
+ ChannelPoolSettings.builder()
+ .setInitialChannelCount(4)
+ .setMinChannelCount(2)
+ .setMaxChannelCount(6)
+ .setMinRpcsPerChannel(1)
+ .setMaxRpcsPerChannel(2)
+ .build(),
+ channelFactory,
+ executor);
+ // With no load, resize() shrinks the pool below its initial channel count.
+ pool.resize();
+ assertThat(pool.entries.get()).hasSize(2);
+
+ pool.refresh();
+ assertThat(pool.entries.get()).hasSize(1);
+
+ // The refill targets the size before the refresh, not the initial channel count.
+ 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(7)).createSingleChannel();
+ }
+
+ @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());
+
+ // Checked and unchecked failures are both handled by expand(); the pool stays usable at its
+ // reduced size.
+ 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);
+ } 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);
+ }
+
+ @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 {
+ 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();
+
+ Mockito.verify(executor).execute(Mockito.any(Runnable.class));
+ 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
+ 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.withSettings().withoutAnnotations());
+
+ 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);
+ ManagedChannel channel2 = mock(ManagedChannel.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);
+ 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 refreshAll_doesNotIncrementGeneration() throws IOException {
+ ManagedChannel channel1 = mock(ManagedChannel.class);
+ ManagedChannel channel2 = mock(ManagedChannel.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);
+ assertThat(pool.getGeneration()).isEqualTo(0);
+
+ // 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 {
+ return createPreemptiveRefreshMtlsPool(size, channelFactory, mockExecutor());
+ }
+
+ private Runnable createPreemptiveRefreshMtlsPool(
+ int size, ChannelFactory channelFactory, ScheduledExecutorService executor)
+ throws IOException {
+ List refreshTasks = new ArrayList<>();
+ 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_whenCertUnreadable_skipsRefreshAndLogs() throws IOException {
+ ManagedChannel initial = Mockito.mock(ManagedChannel.class);
+ ChannelFactory channelFactory = mockChannelFactory();
+ Mockito.when(channelFactory.createSingleChannel()).thenReturn(initial);
+ Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(1, channelFactory);
+
+ // An empty certificate file is what a reader sees while the certificate is being rewritten.
+ java.nio.file.Files.write(tempCert, new byte[0]);
+ 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);
+ }
+
+ Mockito.verify(channelFactory, Mockito.times(1)).createSingleChannel();
+ Mockito.verify(initial, Mockito.never()).shutdown();
+ assertThat(pool.getGeneration()).isEqualTo(0);
+ assertThat(String.join("\n", logHandler.getAllMessages()))
+ .contains("Skipping pre-emptive channel refresh");
+ assertThat(logHandler.getAllMessages()).doesNotContain("Refreshing all channels");
+ }
+
+ @Test
+ void preemptiveRefresh_partialFailureOnCertChange_dropsOldChannelsAndCompletesSwitch()
+ throws IOException {
+ ScheduledExecutorService executor = mockExecutor();
+ ManagedChannel initial1 = Mockito.mock(ManagedChannel.class);
+ ManagedChannel initial2 = Mockito.mock(ManagedChannel.class);
+ ManagedChannel rotated1 = Mockito.mock(ManagedChannel.class);
+ 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(refilled);
+ Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(2, channelFactory, executor);
+
+ pool.invalidateDiskFingerprintCache();
+ writeCert("root_cert.pem");
+ preemptiveRefresh.run();
+
+ // The channel that failed to refresh still uses the old certificate, so it is dropped and
+ // every channel left in the pool uses the new certificate.
+ Mockito.verify(initial1).shutdown();
+ Mockito.verify(initial2).shutdown();
+ assertThat(pool.entries.get()).hasSize(1);
+ assertThat(pool.getGeneration()).isEqualTo(1);
+ pool.invalidateDiskFingerprintCache();
+ assertThat(pool.shouldRefresh()).isFalse();
+
+ // The pool is refilled to its size before the refresh.
+ ArgumentCaptor refillTask = ArgumentCaptor.forClass(Runnable.class);
+ Mockito.verify(executor).execute(refillTask.capture());
+ refillTask.getValue().run();
+ assertThat(pool.entries.get()).hasSize(2);
+ Mockito.verify(refilled, Mockito.never()).shutdown();
+ }
+
+ @Test
+ void preemptiveRefresh_allChannelsFailOnCertChange_leavesRotationPending() throws IOException {
+ ScheduledExecutorService executor = mockExecutor();
+ ManagedChannel initial1 = Mockito.mock(ManagedChannel.class);
+ ManagedChannel initial2 = Mockito.mock(ManagedChannel.class);
+ ChannelFactory channelFactory = mockChannelFactory();
+ Mockito.when(channelFactory.createSingleChannel())
+ .thenReturn(initial1, initial2)
+ .thenThrow(new IOException("Transient failure"));
+ Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(2, channelFactory, executor);
+ List entriesBefore = pool.entries.get();
+
+ pool.invalidateDiskFingerprintCache();
+ writeCert("root_cert.pem");
+ preemptiveRefresh.run();
+
+ // No channel could be recreated, so the pool is left unchanged and the rotation stays pending
+ // for the next periodic or reactive refresh.
+ assertThat(pool.entries.get()).containsExactlyElementsIn(entriesBefore).inOrder();
+ Mockito.verify(initial1, Mockito.never()).shutdown();
+ Mockito.verify(initial2, Mockito.never()).shutdown();
+ assertThat(pool.getGeneration()).isEqualTo(0);
+ pool.invalidateDiskFingerprintCache();
+ assertThat(pool.shouldRefresh()).isTrue();
+ Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class));
+ }
+
+ @Test
+ void preemptiveRefresh_partialFailureWithoutCertChange_keepsOldChannel() throws IOException {
+ 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"));
+ Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(2, channelFactory, executor);
+ ChannelPool.Entry initialEntry2 = pool.entries.get().get(1);
+
+ pool.invalidateDiskFingerprintCache();
+ preemptiveRefresh.run();
+
+ // The certificate did not change, so the channel that failed to refresh keeps its slot.
+ Mockito.verify(initial1).shutdown();
+ Mockito.verify(initial2, Mockito.never()).shutdown();
+ assertThat(pool.entries.get()).hasSize(2);
+ assertThat(pool.entries.get().get(1)).isSameInstanceAs(initialEntry2);
+ assertThat(pool.getGeneration()).isEqualTo(0);
+ Mockito.verify(executor, Mockito.never()).execute(Mockito.any(Runnable.class));
+ }
+
+ /** Waits until {@code thread} is blocked waiting for a monitor lock. */
+ private static void awaitBlocked(Thread thread) throws InterruptedException {
+ long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(5);
+ while (thread.getState() != Thread.State.BLOCKED) {
+ assertThat(System.nanoTime()).isLessThan(deadline);
+ Thread.sleep(1);
+ }
+ }
+
+ /**
+ * Starts a reactive refresh that holds the pool lock while it recreates the channel on the
+ * rotated certificate, runs {@code waiter} on another thread until it blocks on the lock, and
+ * rotates the certificate again before letting the first refresh complete. A waiter that re-read
+ * the certificate would see the second rotation and refresh again.
+ */
+ private void runWhileConcurrentSwitchHoldsLock(
+ CountDownLatch switchStarted, CountDownLatch releaseSwitch, Runnable waiter)
+ throws Exception {
+ Thread switchingThread = new Thread(pool::refresh);
+ Thread waitingThread = new Thread(waiter);
+ try {
+ switchingThread.start();
+ assertThat(switchStarted.await(5, TimeUnit.SECONDS)).isTrue();
+ waitingThread.start();
+ awaitBlocked(waitingThread);
+ writeCert("client_cert.pem");
+ } finally {
+ releaseSwitch.countDown();
+ }
+ switchingThread.join(5000);
+ waitingThread.join(5000);
+ assertThat(switchingThread.isAlive()).isFalse();
+ assertThat(waitingThread.isAlive()).isFalse();
+ }
+
+ @Test
+ void refresh_concurrentSwitchWhileWaitingForLock_skipsWithoutReadingDisk() throws Exception {
+ CountDownLatch switchStarted = new CountDownLatch(1);
+ CountDownLatch releaseSwitch = new CountDownLatch(1);
+ ChannelFactory channelFactory = mockChannelFactory();
+ Mockito.when(channelFactory.createSingleChannel())
+ .thenReturn(Mockito.mock(ManagedChannel.class))
+ .thenAnswer(
+ invocation -> {
+ switchStarted.countDown();
+ releaseSwitch.await();
+ return Mockito.mock(ManagedChannel.class);
+ })
+ .thenReturn(Mockito.mock(ManagedChannel.class));
+ createMtlsPoolAndRotateCert(
+ ChannelPoolSettings.staticallySized(1), channelFactory, mockExecutor());
+
+ runWhileConcurrentSwitchHoldsLock(switchStarted, releaseSwitch, pool::refresh);
+
+ // Only the initial channel and the concurrent switch created channels; the waiter skipped.
+ Mockito.verify(channelFactory, Mockito.times(2)).createSingleChannel();
+ assertThat(pool.getGeneration()).isEqualTo(1);
+ // The second rotation is still detected for the next UNAUTHENTICATED failure.
+ pool.invalidateDiskFingerprintCache();
+ assertThat(pool.shouldRefresh()).isTrue();
+ }
+
+ @Test
+ void preemptiveRefresh_concurrentSwitchWhileWaitingForLock_skipsRefresh() throws Exception {
+ CountDownLatch switchStarted = new CountDownLatch(1);
+ CountDownLatch releaseSwitch = new CountDownLatch(1);
+ ChannelFactory channelFactory = mockChannelFactory();
+ Mockito.when(channelFactory.createSingleChannel())
+ .thenReturn(Mockito.mock(ManagedChannel.class))
+ .thenAnswer(
+ invocation -> {
+ switchStarted.countDown();
+ releaseSwitch.await();
+ return Mockito.mock(ManagedChannel.class);
+ })
+ .thenReturn(Mockito.mock(ManagedChannel.class));
+ Runnable preemptiveRefresh = createPreemptiveRefreshMtlsPool(1, channelFactory);
+ pool.invalidateDiskFingerprintCache();
+ writeCert("root_cert.pem");
+
+ runWhileConcurrentSwitchHoldsLock(switchStarted, releaseSwitch, preemptiveRefresh);
+
+ // The channels were just recreated by the concurrent switch, so the periodic refresh skipped.
+ Mockito.verify(channelFactory, Mockito.times(2)).createSingleChannel();
+ assertThat(pool.getGeneration()).isEqualTo(1);
+ pool.invalidateDiskFingerprintCache();
+ assertThat(pool.shouldRefresh()).isTrue();
+ }
+
+ @Test
+ void refresh_whenCertUnchanged_noOpsAndDoesNotIncrementGeneration() 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
void channelRefreshShouldSwapChannels() throws IOException {
ManagedChannel underlyingChannel1 = mock(ManagedChannel.class);
@@ -450,7 +1432,8 @@ void channelRefreshShouldSwapChannels() throws IOException {
.setPreemptiveRefreshEnabled(true)
.build(),
channelFactory,
- provider);
+ provider,
+ null);
Mockito.reset(underlyingChannel1);
pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT);
@@ -459,10 +1442,41 @@ void channelRefreshShouldSwapChannels() throws IOException {
.newCall(Mockito.>any(), Mockito.any(CallOptions.class));
// swap channel
- pool.refresh();
+ pool.refreshAll();
+
+ pool.newCall(FakeMethodDescriptor.create(), CallOptions.DEFAULT);
+
+ Mockito.verify(underlyingChannel2, Mockito.only())
+ .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));
}
@@ -486,7 +1500,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 +1568,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 +1602,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 +1631,8 @@ void removedActiveChannelsAreShutdown() throws Exception {
.setMaxRpcsPerChannel(2)
.build(),
channelFactory,
- provider);
+ provider,
+ null);
assertThat(pool.entries.get()).hasSize(2);
// Start 2 RPCs
@@ -652,7 +1670,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 +1735,8 @@ void repeatedResizingLogsWarningOnExpand() throws Exception {
.setMaxChannelCount(10)
.build(),
channelFactory,
- provider);
+ provider,
+ null);
assertThat(pool.entries.get()).hasSize(1);
FakeLogHandler logHandler = new FakeLogHandler();
@@ -769,7 +1788,8 @@ void repeatedResizingLogsWarningOnShrink() throws Exception {
.setMaxChannelCount(10)
.build(),
channelFactory,
- provider);
+ provider,
+ null);
assertThat(pool.entries.get()).hasSize(10);
FakeLogHandler logHandler = new FakeLogHandler();
@@ -805,7 +1825,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 +1863,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 +1900,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 +1936,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
@@ -924,4 +1947,278 @@ 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.withSettings().withoutAnnotations());
+ 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.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"));
+
+ 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.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"))
+ .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.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"))
+ .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.withSettings().withoutAnnotations());
+ 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();
+ }
+
+ // 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();
+ 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 e20767fdb8ed..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);
@@ -494,4 +510,64 @@ 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);
+ }
+
+ @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());
+
+ // 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/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-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..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
@@ -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;
@@ -664,7 +665,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 +685,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 +716,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 +883,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 +899,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 +1211,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 +1258,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);
}
@@ -1342,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_whenMtlsActiveAndKeyStoreNull_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 2679b51860df..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
@@ -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,12 @@ public HttpJsonCallContext merge(ApiCallContext inputCallContext) {
newRetryableCodes = this.retryableCodes;
}
+ TransportChannel newTransportChannel = httpJsonCallContext.transportChannel;
+ if (newTransportChannel == null
+ && (httpJsonCallContext.channel == null || httpJsonCallContext.channel.equals(channel))) {
+ 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 +244,8 @@ public HttpJsonCallContext merge(ApiCallContext inputCallContext) {
newTracer,
newRetrySettings,
newRetryableCodes,
- endpointContext);
+ endpointContext,
+ newTransportChannel);
}
@Override
@@ -251,7 +263,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 +304,8 @@ public HttpJsonCallContext withEndpointContext(EndpointContext endpointContext)
this.tracer,
this.retrySettings,
this.retryableCodes,
- endpointContext);
+ endpointContext,
+ this.transportChannel);
}
@Override
@@ -301,7 +331,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 +377,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 +428,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 +466,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 +491,8 @@ public ApiCallContext withOption(Key key, T value) {
this.tracer,
this.retrySettings,
this.retryableCodes,
- this.endpointContext);
+ this.endpointContext,
+ this.transportChannel);
}
/** {@inheritDoc} */
@@ -527,7 +562,8 @@ public HttpJsonCallContext withRetrySettings(RetrySettings retrySettings) {
this.tracer,
retrySettings,
this.retryableCodes,
- this.endpointContext);
+ this.endpointContext,
+ this.transportChannel);
}
@Override
@@ -548,7 +584,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 +600,8 @@ public HttpJsonCallContext withChannel(@Nullable HttpJsonChannel newChannel) {
this.tracer,
this.retrySettings,
this.retryableCodes,
- this.endpointContext);
+ this.endpointContext,
+ (newChannel != null && newChannel.equals(this.channel)) ? this.transportChannel : null);
}
public HttpJsonCallContext withCallOptions(HttpJsonCallOptions newCallOptions) {
@@ -578,7 +616,8 @@ public HttpJsonCallContext withCallOptions(HttpJsonCallOptions newCallOptions) {
this.tracer,
this.retrySettings,
this.retryableCodes,
- this.endpointContext);
+ this.endpointContext,
+ this.transportChannel);
}
@Deprecated
@@ -614,7 +653,8 @@ public HttpJsonCallContext withTracer(@Nonnull ApiTracer newTracer) {
newTracer,
this.retrySettings,
this.retryableCodes,
- this.endpointContext);
+ this.endpointContext,
+ this.transportChannel);
}
@Override
@@ -634,7 +674,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 +689,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..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
@@ -64,6 +64,21 @@ public HttpJsonChannel getChannel() {
return getManagedChannel();
}
+ @Override
+ public void refresh() {
+ getManagedChannel().refresh();
+ }
+
+ @Override
+ 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 92ce4efe36aa..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
@@ -233,33 +233,90 @@ private NetHttpTransport.Builder configureMtls(NetHttpTransport.Builder builder)
return builder;
}
- private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecurityException {
- HttpTransport httpTransportToUse = httpTransport;
- if (httpTransportToUse == null) {
- httpTransportToUse = createHttpTransport();
+ 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 {
+ return buildManagedChannel(
+ httpTransport != null ? httpTransport : createChannelHttpTransport());
+ }
- // 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();
-
- 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);
+ private ManagedHttpJsonChannel buildManagedChannel(HttpTransport httpTransportToUse) {
+ return ManagedHttpJsonChannel.newBuilder()
+ .setEndpoint(endpoint)
+ .setExecutor(executor)
+ .setHttpTransport(httpTransportToUse)
+ .setManageHttpTransport(httpTransport == null)
+ .build();
+ }
+
+ private HttpJsonTransportChannel createChannel() throws IOException, GeneralSecurityException {
+ boolean isMtlsActive =
+ httpTransport == null
+ && mtlsProvider != null
+ && certificateBasedAccess.useMtlsClientCertificate();
+ String workloadCertPath = isMtlsActive ? certificateBasedAccess.getWorkloadCertPath() : null;
+
+ ManagedHttpJsonChannel baseChannel;
+ if (workloadCertPath != null) {
+ java.util.function.Supplier transportFactory =
+ () -> {
+ try {
+ return createChannelHttpTransport();
+ } catch (IOException | GeneralSecurityException e) {
+ throw new java.lang.RuntimeException("Failed to create mTLS HttpTransport", e);
+ }
+ };
+ // RefreshingHttpJsonChannel records the baseline certificate fingerprint before creating the
+ // initial transport, so a rotation during startup is detected on the next auth failure.
+ try {
+ baseChannel =
+ new RefreshingHttpJsonChannel(
+ transportFactory, this::buildManagedChannel, 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();
}
- return HttpJsonTransportChannel.newBuilder().setManagedChannel(channel).build();
+ try {
+ ManagedHttpJsonChannel channel = baseChannel;
+
+ 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();
+ } catch (Throwable t) {
+ baseChannel.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/ManagedHttpJsonChannel.java b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/ManagedHttpJsonChannel.java
index 87767bee5c7f..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,12 +51,35 @@ 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;
protected ManagedHttpJsonChannel() {
- this(null, true, null, null);
+ this(null, true, null, null, true);
+ }
+
+ /**
+ * 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;
+ this.httpTransport = null;
+ this.usingDefaultTransport = false;
+ 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;
}
String getEndpoint() {
@@ -68,11 +91,20 @@ 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,
@Nullable String endpoint,
- @Nullable HttpTransport httpTransport) {
+ @Nullable HttpTransport httpTransport,
+ boolean usingDefaultTransport) {
this.executor = executor;
this.usingDefaultExecutor = usingDefaultExecutor;
this.endpoint = endpoint;
@@ -82,6 +114,7 @@ private ManagedHttpJsonChannel(
new NetHttpTransport.Builder())
.build()
: httpTransport;
+ this.usingDefaultTransport = usingDefaultTransport || httpTransport == null;
this.deadlineScheduledExecutorService = Executors.newSingleThreadScheduledExecutor();
}
@@ -98,6 +131,22 @@ public HttpJsonClientCall newCall(
deadlineScheduledExecutorService);
}
+ /**
+ * 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;
+ }
+
@VisibleForTesting
Executor getExecutor() {
return executor;
@@ -116,7 +165,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 +209,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 +258,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 +280,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 +297,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..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
@@ -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;
@@ -41,22 +43,57 @@ class ManagedHttpJsonInterceptorChannel extends ManagedHttpJsonChannel {
ManagedHttpJsonInterceptorChannel(
ManagedHttpJsonChannel channel, HttpJsonClientInterceptor interceptor) {
- super();
+ super(true);
this.channel = channel;
this.interceptor = interceptor;
}
+ /** {@inheritDoc} */
+ @Override
+ public long getGeneration() {
+ return channel.getGeneration();
+ }
+
@VisibleForTesting
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);
}
+ /** {@inheritDoc} */
+ @Override
+ public void refresh() {
+ channel.refresh();
+ }
+
+ /** {@inheritDoc} */
+ @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..6d20a69139f8
--- /dev/null
+++ b/sdk-platform-java/gax-java/gax-httpjson/src/main/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannel.java
@@ -0,0 +1,213 @@
+/*
+ * 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.rpc.mtls.CertificateRotationTracker;
+import com.google.api.gax.rpc.mtls.WorkloadCertificateUtils;
+import com.google.common.annotations.VisibleForTesting;
+import java.util.concurrent.Executor;
+import java.util.concurrent.TimeUnit;
+import java.util.concurrent.atomic.AtomicLong;
+import java.util.function.Function;
+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. 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 {
+
+ private static final Logger LOG = Logger.getLogger(RefreshingHttpJsonChannel.class.getName());
+
+ private final CertificateRotationTracker rotationTracker;
+ private final Supplier transportFactory;
+ private final String workloadCertPath;
+ 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 transportFactory,
+ Function channelFactory,
+ String workloadCertPath) {
+ super(true);
+ this.transportFactory = transportFactory;
+ this.workloadCertPath = workloadCertPath;
+ // 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);
+ this.delegate = channelFactory.apply(transportFactory.get());
+ }
+
+ @VisibleForTesting
+ String getWorkloadCertPath() {
+ return workloadCertPath;
+ }
+
+ @VisibleForTesting
+ String getCertificateFingerprint(String certPath) {
+ return WorkloadCertificateUtils.getCertificateFingerprint(certPath);
+ }
+
+ /** {@inheritDoc} */
+ @Override
+ public boolean shouldRefresh() {
+ return rotationTracker.shouldRefresh();
+ }
+
+ /** {@inheritDoc} */
+ @Override
+ public void refresh() {
+ // A generation change while waiting for the lock means a concurrent refresh already swapped in
+ // a transport with a new certificate, so there is no need to read the certificate again.
+ long generationBeforeLock = generation.get();
+ synchronized (refreshLock) {
+ if (isShutdown()) {
+ return;
+ }
+ if (generation.get() != generationBeforeLock) {
+ LOG.fine(
+ "HTTP/JSON channel was already refreshed by a concurrent thread, skipping duplicate"
+ + " refresh");
+ return;
+ }
+ String currentDiskFingerprint = rotationTracker.readDiskFingerprint();
+ if (currentDiskFingerprint.isEmpty()) {
+ return;
+ }
+
+ // Double-check inside refreshLock
+ if (rotationTracker.isAlreadyActive(currentDiskFingerprint)) {
+ LOG.fine(
+ "HTTP/JSON channel was already refreshed by a concurrent thread, skipping duplicate"
+ + " refresh");
+ return;
+ }
+
+ LOG.info("mTLS certificate rotation detected. Refreshing HTTP/JSON transport.");
+
+ HttpTransport newTransport;
+ try {
+ newTransport = transportFactory.get();
+ } catch (Exception e) {
+ LOG.log(Level.WARNING, "Failed to refresh HTTP/JSON transport, keeping old transport", e);
+ return;
+ }
+ // The previous transport is not shut down: calls created before the swap may still be using
+ // 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.
+ generation.incrementAndGet();
+ rotationTracker.markRefreshed(currentDiskFingerprint);
+ }
+ }
+
+ /** {@inheritDoc} */
+ @Override
+ public long getGeneration() {
+ return generation.get();
+ }
+
+ @Override
+ public HttpJsonClientCall newCall(
+ ApiMethodDescriptor methodDescriptor, HttpJsonCallOptions callOptions) {
+ return delegate.newCall(methodDescriptor, callOptions);
+ }
+
+ @Override
+ Executor getExecutor() {
+ return delegate.getExecutor();
+ }
+
+ @Override
+ String getEndpoint() {
+ return delegate.getEndpoint();
+ }
+
+ @Override
+ @VisibleForTesting
+ HttpTransport getHttpTransport() {
+ return delegate.getHttpTransport();
+ }
+
+ @VisibleForTesting
+ void invalidateDiskFingerprintCache() {
+ rotationTracker.invalidateCache();
+ }
+
+ @Override
+ public void shutdown() {
+ // Serialized with refresh() so that a transport swap cannot race with shutdown.
+ synchronized (refreshLock) {
+ delegate.shutdown();
+ }
+ }
+
+ @Override
+ public boolean isShutdown() {
+ return delegate.isShutdown();
+ }
+
+ @Override
+ public boolean isTerminated() {
+ return delegate.isTerminated();
+ }
+
+ @Override
+ public void shutdownNow() {
+ synchronized (refreshLock) {
+ delegate.shutdownNow();
+ }
+ }
+
+ @Override
+ public boolean awaitTermination(long duration, TimeUnit unit) throws InterruptedException {
+ return delegate.awaitTermination(duration, unit);
+ }
+
+ @Override
+ public void close() {
+ shutdown();
+ }
+}
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..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
@@ -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);
@@ -334,4 +348,55 @@ 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
+ Truth.assertThat(context.withChannel(channel1).getTransportChannel())
+ .isSameInstanceAs(transportChannel1);
+
+ // 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 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())
+ .isSameInstanceAs(transportChannel1);
+ }
+
+ @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 8c95c1d2e1c4..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
@@ -31,13 +31,16 @@
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;
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;
@@ -46,6 +49,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 +59,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,20 +184,361 @@ void managedChannelUsesCustomExecutor() throws IOException {
instantiatingHttpJsonChannelProvider.getTransportChannel().shutdownNow();
}
- @Override
- protected Object getMtlsObjectFromTransportChannel(
- MtlsProvider provider, CertificateBasedAccess certificateBasedAccess)
+ @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 (direct ManagedHttpJsonChannel when workloadCertPath is
+ // null)
+ ManagedHttpJsonInterceptorChannel interceptorChannel =
+ (ManagedHttpJsonInterceptorChannel) httpJsonTransportChannel.getManagedChannel();
+ ManagedHttpJsonInterceptorChannel managedHttpJsonChannel =
+ (ManagedHttpJsonInterceptorChannel) interceptorChannel.getChannel();
+ ManagedHttpJsonChannel channel = managedHttpJsonChannel.getChannel();
+
+ 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();
+ }
+
+ @Test
+ void channelCreation_withWorkloadCertPath_wrapsWithRefreshingHttpJsonChannel()
+ 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);
+
+ 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 channelCreation_withCustomHttpTransport_ignoresWorkloadCertPathAndDoesNotWrap()
+ 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);
+
+ InstantiatingHttpJsonChannelProvider provider =
+ InstantiatingHttpJsonChannelProvider.newBuilder()
+ .setEndpoint(DEFAULT_ENDPOINT)
+ .setHttpTransport(mockHttpTransport)
+ .setMtlsProvider(mtlsProvider)
+ .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_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 {
+ Mockito.when(certificateBasedAccess.useMtlsClientCertificate()).thenReturn(true);
+ Mockito.when(certificateBasedAccess.getWorkloadCertPath()).thenReturn("fake/cert/path.json");
+ MtlsProvider failingMtlsProvider =
+ Mockito.mock(MtlsProvider.class, Mockito.withSettings().withoutAnnotations());
+ 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 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_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.
+ 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_withMtlsKeyStore_returnsMtlsTransport()
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("localhost:8080")
+ .setEndpoint(DEFAULT_ENDPOINT)
.setMtlsProvider(provider)
.setCertificateBasedAccess(certificateBasedAccess)
- .setHeaderProvider(Collections::emptyMap)
- .setExecutor(Runnable::run)
.build();
- NetHttpTransport transport = (NetHttpTransport) channelProvider.createHttpTransport();
- return (transport != null && transport.isMtls()) ? transport : null;
+
+ 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
+ void createHttpTransport_whenMtlsProviderNullOrNotUsingClientCert_returnsNonMtlsTransport()
+ throws IOException, GeneralSecurityException {
+ InstantiatingHttpJsonChannelProvider nullMtlsProviderChannelProvider =
+ InstantiatingHttpJsonChannelProvider.newBuilder()
+ .setEndpoint(DEFAULT_ENDPOINT)
+ .setMtlsProvider(null)
+ .setCertificateBasedAccess(certificateBasedAccess)
+ .build();
+ NetHttpTransport nullMtlsProviderTransport =
+ (NetHttpTransport) nullMtlsProviderChannelProvider.createHttpTransport();
+ assertThat(nullMtlsProviderTransport).isNotNull();
+ assertThat(nullMtlsProviderTransport.isMtls()).isFalse();
+
+ 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();
+ NetHttpTransport disabledMtlsTransport =
+ (NetHttpTransport) disabledMtlsChannelProvider.createHttpTransport();
+ assertThat(disabledMtlsTransport).isNotNull();
+ assertThat(disabledMtlsTransport.isMtls()).isFalse();
}
@Test
@@ -214,4 +560,21 @@ void testConfigureConscryptSecurityProvider_returnsConfiguredBuilder() {
HttpJsonConscryptUtils.configureConscryptSecurityProvider(builder);
assertThat(result).isSameInstanceAs(builder);
}
+
+ @Override
+ protected Object getMtlsObjectFromTransportChannel(
+ MtlsProvider provider, CertificateBasedAccess certificateBasedAccess)
+ throws IOException, GeneralSecurityException {
+ InstantiatingHttpJsonChannelProvider channelProvider =
+ InstantiatingHttpJsonChannelProvider.newBuilder()
+ .setEndpoint("localhost:8080")
+ .setMtlsProvider(provider)
+ .setCertificateBasedAccess(certificateBasedAccess)
+ .setHeaderProvider(
+ mock(HeaderProvider.class, Mockito.withSettings().withoutAnnotations()))
+ .setExecutor(mock(Executor.class))
+ .build();
+ 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
new file mode 100644
index 000000000000..1018005f2df5
--- /dev/null
+++ b/sdk-platform-java/gax-java/gax-httpjson/src/test/java/com/google/api/gax/httpjson/RefreshingHttpJsonChannelTest.java
@@ -0,0 +1,607 @@
+/*
+ * 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.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.CountDownLatch;
+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;
+import org.junit.jupiter.api.BeforeEach;
+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 {
+ @Override
+ public void start(Listener responseListener, HttpJsonMetadata requestHeaders) {}
+
+ @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 transportFactoryCount;
+ private HttpTransport lastCreatedTransport;
+ private FakeManagedHttpJsonChannel lastCreatedChannel;
+ private String testCertPath;
+ private String testFingerprint;
+ private boolean shouldThrowOnFactory;
+ private List createdChannels;
+
+ private Supplier transportFactory =
+ () -> {
+ if (shouldThrowOnFactory) {
+ throw new RuntimeException("Simulated factory failure");
+ }
+ transportFactoryCount.incrementAndGet();
+ lastCreatedTransport = new MockHttpTransport();
+ return lastCreatedTransport;
+ };
+
+ private Function channelFactory =
+ transport -> {
+ lastCreatedChannel = new FakeManagedHttpJsonChannel();
+ lastCreatedChannel.setHttpTransport(transport);
+ return lastCreatedChannel;
+ };
+
+ @BeforeEach
+ void setUp() {
+ transportFactoryCount = 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(
+ () -> transportFactory.get(),
+ transport -> channelFactory.apply(transport),
+ "fake/cert/path.json") {
+ @Override
+ String getWorkloadCertPath() {
+ return testCertPath;
+ }
+
+ @Override
+ String getCertificateFingerprint(String certPath) {
+ return testFingerprint;
+ }
+ };
+ createdChannels.add(ch);
+ return ch;
+ }
+
+ private void rotateCertificate(RefreshingHttpJsonChannel channel) {
+ channel.invalidateDiskFingerprintCache();
+ testFingerprint = "fingerprint2";
+ }
+
+ @Test
+ void testShouldRefreshFalseWhenCertPathNull() {
+ testCertPath = null;
+ RefreshingHttpJsonChannel channel = createTestChannel();
+ assertFalse(channel.shouldRefresh());
+ }
+
+ @Test
+ void testShouldRefreshFalseWhenUnchanged() {
+ RefreshingHttpJsonChannel channel = createTestChannel();
+
+ channel.invalidateDiskFingerprintCache(); // Invalidate 1-second cache
+ assertFalse(channel.shouldRefresh());
+ }
+
+ @Test
+ void testShouldRefreshTrueWhenChanged() {
+ RefreshingHttpJsonChannel channel = createTestChannel();
+
+ rotateCertificate(channel);
+
+ 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 rotationDuringInitialTransportCreation_isDetectedAndRefreshed() {
+ transportFactory =
+ () -> {
+ if (transportFactoryCount.incrementAndGet() == 1) {
+ // Simulate the certificate rotating on disk while the initial transport loads it.
+ testFingerprint = "fingerprint2";
+ }
+ lastCreatedTransport = new MockHttpTransport();
+ return lastCreatedTransport;
+ };
+
+ RefreshingHttpJsonChannel channel = createTestChannel();
+
+ // The baseline was recorded before the initial transport was created, so the rotation is seen.
+ assertTrue(channel.shouldRefresh());
+
+ channel.refresh();
+ assertEquals(2, transportFactoryCount.get());
+ assertEquals(1, channel.getGeneration());
+ assertFalse(channel.shouldRefresh());
+ }
+
+ @Test
+ void refresh_swapsTransportAndKeepsChannel() {
+ RefreshingHttpJsonChannel channel = createTestChannel();
+ FakeManagedHttpJsonChannel underlyingChannel = lastCreatedChannel;
+ HttpTransport initialTransport = channel.getHttpTransport();
+ assertSame(lastCreatedTransport, initialTransport);
+
+ rotateCertificate(channel);
+ channel.refresh();
+
+ 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 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 =
+ 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();
+ HttpJsonCallOptions callOptions = HttpJsonCallOptions.newBuilder().build();
+ HttpJsonCallContext callContext = HttpJsonCallContext.createDefault();
+
+ HttpJsonClientCall callBeforeRefresh =
+ channel.newCall(FAKE_METHOD_DESCRIPTOR, callOptions);
+
+ rotateCertificate(channel);
+ channel.refresh();
+ assertEquals(1, channel.getGeneration());
+
+ 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 testRefreshDoesNotCreateTransportWhenShutdown() {
+ RefreshingHttpJsonChannel channel = createTestChannel();
+ assertEquals(1, transportFactoryCount.get());
+
+ channel.shutdown();
+ rotateCertificate(channel);
+ channel.refresh();
+
+ assertEquals(1, transportFactoryCount.get());
+ 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 refresh_concurrentRefreshWhileWaitingForLock_skipsWithoutReadingDisk() throws Exception {
+ CountDownLatch refreshStarted = new CountDownLatch(1);
+ CountDownLatch releaseRefresh = new CountDownLatch(1);
+ transportFactory =
+ () -> {
+ if (transportFactoryCount.incrementAndGet() == 2) {
+ refreshStarted.countDown();
+ try {
+ releaseRefresh.await();
+ } catch (InterruptedException e) {
+ Thread.currentThread().interrupt();
+ }
+ }
+ return new MockHttpTransport();
+ };
+ RefreshingHttpJsonChannel channel = createTestChannel();
+ rotateCertificate(channel);
+
+ Thread refreshingThread = new Thread(channel::refresh);
+ Thread waitingThread = new Thread(channel::refresh);
+ try {
+ refreshingThread.start();
+ assertTrue(refreshStarted.await(5, TimeUnit.SECONDS));
+ waitingThread.start();
+ long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(5);
+ while (waitingThread.getState() != Thread.State.BLOCKED) {
+ assertTrue(System.nanoTime() < deadline);
+ Thread.sleep(1);
+ }
+ // Rotate again before the first refresh completes. A waiter that re-read the certificate
+ // would see this fingerprint and refresh a second time.
+ testFingerprint = "fingerprint3";
+ } finally {
+ releaseRefresh.countDown();
+ }
+ refreshingThread.join(5000);
+ waitingThread.join(5000);
+ assertFalse(refreshingThread.isAlive());
+ assertFalse(waitingThread.isAlive());
+
+ // Only the initial transport and the concurrent refresh created transports; the waiter skipped.
+ assertEquals(2, transportFactoryCount.get());
+ assertEquals(1, channel.getGeneration());
+ // The second rotation is still detected for the next UNAUTHENTICATED failure.
+ channel.invalidateDiskFingerprintCache();
+ assertTrue(channel.shouldRefresh());
+ }
+
+ @Test
+ void testRefreshFactoryExceptionDoesNotWedgeFingerprint() {
+ RefreshingHttpJsonChannel channel = createTestChannel();
+ HttpTransport initialTransport = channel.getHttpTransport();
+ assertEquals(1, transportFactoryCount.get());
+
+ shouldThrowOnFactory = true;
+ rotateCertificate(channel);
+
+ // Factory failure is logged and the existing transport is kept
+ channel.refresh();
+ assertEquals(1, transportFactoryCount.get());
+ assertEquals(0, channel.getGeneration());
+ assertSame(initialTransport, channel.getHttpTransport());
+
+ // 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, transportFactoryCount.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();
+ lastCreatedChannel.isTerminated = true;
+
+ channel.shutdown();
+ assertTrue(channel.awaitTermination(0, TimeUnit.MILLISECONDS));
+ }
+
+ @Test
+ void testChannelDelegationMethods() {
+ RefreshingHttpJsonChannel channel = createTestChannel();
+ FakeManagedHttpJsonChannel underlyingChannel = lastCreatedChannel;
+ FakeHttpJsonClientCall