diff --git a/core/src/main/java/io/grpc/internal/ManagedChannelImpl.java b/core/src/main/java/io/grpc/internal/ManagedChannelImpl.java index 00df05a0c00..e533770da22 100644 --- a/core/src/main/java/io/grpc/internal/ManagedChannelImpl.java +++ b/core/src/main/java/io/grpc/internal/ManagedChannelImpl.java @@ -173,7 +173,7 @@ public Result selectConfig(PickSubchannelArgs args) { private final NameResolverProvider nameResolverProvider; private final NameResolver.Args nameResolverArgs; private final LoadBalancerProvider loadBalancerFactory; - private final ClientTransportFactory originalTransportFactory; + private final RefCountedClientTransportFactory originalTransportFactory; @Nullable private final ChannelCredentials originalChannelCreds; private final ClientTransportFactory transportFactory; @@ -562,11 +562,15 @@ ClientStream newSubstream( this.executorPool = checkNotNull(builder.executorPool, "executorPool"); this.executor = checkNotNull(executorPool.getObject(), "executor"); this.originalChannelCreds = builder.channelCredentials; - this.originalTransportFactory = clientTransportFactory; + if (clientTransportFactory instanceof RefCountedClientTransportFactory) { + this.originalTransportFactory = (RefCountedClientTransportFactory) clientTransportFactory; + } else { + this.originalTransportFactory = new RefCountedClientTransportFactory(clientTransportFactory); + } this.offloadExecutorHolder = new ExecutorHolder(checkNotNull(builder.offloadExecutorPool, "offloadExecutorPool")); this.transportFactory = new CallCredentialsApplyingTransportFactory( - clientTransportFactory, builder.callCredentials, this.offloadExecutorHolder); + originalTransportFactory, builder.callCredentials, this.offloadExecutorHolder); this.scheduledExecutor = new RestrictedScheduledExecutor(transportFactory.getScheduledExecutorService()); maxTraceEvents = builder.maxTraceEvents; @@ -1462,7 +1466,10 @@ final class ResolvingOobChannelBuilder final ClientTransportFactory transportFactory; CallCredentials callCredentials; if (channelCreds instanceof DefaultChannelCreds) { - transportFactory = originalTransportFactory; + // TODO(kannanjgithub) We should eventually refactor ManagedChannelImplBuilder so + // callCredentials can be resolved lazily at build() time, allowing transport factory + // retention to happen strictly inside buildClientTransportFactory(). + transportFactory = originalTransportFactory.retain(); callCredentials = null; } else { SwapChannelCredentialsResult swapResult = diff --git a/core/src/main/java/io/grpc/internal/RefCountedClientTransportFactory.java b/core/src/main/java/io/grpc/internal/RefCountedClientTransportFactory.java new file mode 100644 index 00000000000..b1962b1e9f0 --- /dev/null +++ b/core/src/main/java/io/grpc/internal/RefCountedClientTransportFactory.java @@ -0,0 +1,75 @@ +/* + * Copyright 2026 The gRPC Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package io.grpc.internal; + +import static com.google.common.base.Preconditions.checkNotNull; +import static com.google.common.base.Preconditions.checkState; + +import io.grpc.ChannelCredentials; +import io.grpc.ChannelLogger; +import java.net.SocketAddress; +import java.util.Collection; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.atomic.AtomicInteger; + +/** + * A wrapper for {@link ClientTransportFactory} that reference-counts calls to {@link #retain()} and + * {@link #close()}, ensuring the delegate factory is closed only when all references are released. + */ +final class RefCountedClientTransportFactory implements ClientTransportFactory { + private final ClientTransportFactory delegate; + private final AtomicInteger refCount = new AtomicInteger(1); + + public RefCountedClientTransportFactory(ClientTransportFactory delegate) { + this.delegate = checkNotNull(delegate, "delegate"); + } + + public RefCountedClientTransportFactory retain() { + refCount.incrementAndGet(); + return this; + } + + @Override + public ConnectionClientTransport newClientTransport( + SocketAddress serverAddress, ClientTransportOptions options, ChannelLogger channelLogger) { + return delegate.newClientTransport(serverAddress, options, channelLogger); + } + + @Override + public ScheduledExecutorService getScheduledExecutorService() { + return delegate.getScheduledExecutorService(); + } + + @Override + public Collection> getSupportedSocketAddressTypes() { + return delegate.getSupportedSocketAddressTypes(); + } + + @Override + public SwapChannelCredentialsResult swapChannelCredentials(ChannelCredentials channelCreds) { + return delegate.swapChannelCredentials(channelCreds); + } + + @Override + public void close() { + int count = refCount.decrementAndGet(); + checkState(count >= 0, "Reference count has gone negative: %s", count); + if (count == 0) { + delegate.close(); + } + } +} diff --git a/core/src/test/java/io/grpc/internal/ManagedChannelImplTest.java b/core/src/test/java/io/grpc/internal/ManagedChannelImplTest.java index e958fcdae00..d052cce137b 100644 --- a/core/src/test/java/io/grpc/internal/ManagedChannelImplTest.java +++ b/core/src/test/java/io/grpc/internal/ManagedChannelImplTest.java @@ -4824,6 +4824,46 @@ public void run() { }); } + @Test + public void oobChannelTermination_doesNotCloseSharedTransportFactory() { + channelBuilder.nameResolverRegistry.register(new NameResolverProvider() { + @Override + public NameResolver newNameResolver(URI targetUri, NameResolver.Args args) { + NameResolver resolver = mock(NameResolver.class); + when(resolver.getServiceAuthority()).thenReturn( + targetUri.getAuthority() != null ? targetUri.getAuthority() : targetUri.getPath()); + return resolver; + } + + @Override + public String getDefaultScheme() { + return expectedUri.getScheme(); + } + + @Override + protected boolean isAvailable() { + return true; + } + + @Override + protected int priority() { + return 10; + } + }); + createChannel(); + ManagedChannel oob = helper.createResolvingOobChannelBuilder("oobauthority").build(); + + // Shutting down OOB channel should release its reference but not close the + // shared transport factory + oob.shutdownNow(); + verify(mockTransportFactory, never()).close(); + + // Terminating the main channel releases the final reference and closes the + // transport factory + channel.shutdownNow(); + verify(mockTransportFactory).close(); + } + @SuppressWarnings("unchecked") private static Map parseConfig(String json) throws Exception { return (Map) JsonParser.parse(json); diff --git a/core/src/test/java/io/grpc/internal/RefCountedClientTransportFactoryTest.java b/core/src/test/java/io/grpc/internal/RefCountedClientTransportFactoryTest.java new file mode 100644 index 00000000000..3ead6396ce5 --- /dev/null +++ b/core/src/test/java/io/grpc/internal/RefCountedClientTransportFactoryTest.java @@ -0,0 +1,81 @@ +/* + * Copyright 2026 The gRPC Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package io.grpc.internal; + +import static org.junit.Assert.assertThrows; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; + +import org.junit.Rule; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; +import org.mockito.Mock; +import org.mockito.junit.MockitoJUnit; +import org.mockito.junit.MockitoRule; + +/** Unit tests for {@link RefCountedClientTransportFactory}. */ +@RunWith(JUnit4.class) +public class RefCountedClientTransportFactoryTest { + @Rule public final MockitoRule mocks = MockitoJUnit.rule(); + + @Mock private ClientTransportFactory mockDelegate; + + @Test + public void singleClose_closesDelegate() { + RefCountedClientTransportFactory factory = new RefCountedClientTransportFactory(mockDelegate); + factory.close(); + verify(mockDelegate).close(); + } + + @Test + public void retainAndClose_closesDelegateOnlyWhenCountReachesZero() { + RefCountedClientTransportFactory factory = new RefCountedClientTransportFactory(mockDelegate); + RefCountedClientTransportFactory retained = factory.retain(); + + factory.close(); + verify(mockDelegate, never()).close(); + + retained.close(); + verify(mockDelegate).close(); + } + + @Test + public void multipleRetains_requiresEqualClosesToCloseDelegate() { + RefCountedClientTransportFactory factory = new RefCountedClientTransportFactory(mockDelegate); + factory.retain(); + factory.retain(); + + factory.close(); + verify(mockDelegate, never()).close(); + + factory.close(); + verify(mockDelegate, never()).close(); + + factory.close(); + verify(mockDelegate).close(); + } + + @Test + public void closeMoreThanRetain_throwsIllegalStateException() { + RefCountedClientTransportFactory factory = new RefCountedClientTransportFactory(mockDelegate); + factory.close(); + verify(mockDelegate).close(); + + assertThrows(IllegalStateException.class, factory::close); + } +}