Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -19,11 +19,11 @@
import static com.google.common.base.Preconditions.checkNotNull;

import com.google.common.annotations.VisibleForTesting;
import com.google.common.collect.ImmutableSet;
import io.grpc.Status;
import io.grpc.xds.internal.security.Closeable;
import java.security.PrivateKey;
import java.security.cert.X509Certificate;
import java.util.Collections;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
Expand Down Expand Up @@ -78,8 +78,9 @@ synchronized void removeWatcher(Watcher watcher) {
downstreamWatchers.remove(watcher);
}

@VisibleForTesting public Set<Watcher> getDownstreamWatchers() {
return Collections.unmodifiableSet(downstreamWatchers);
@VisibleForTesting public synchronized Set<Watcher> getDownstreamWatchers() {
// Callbacks may register or remove watchers reentrantly, even while we hold the lock.
return ImmutableSet.copyOf(downstreamWatchers);
}

private void sendLastCertificateUpdate(Watcher watcher) {
Expand All @@ -100,7 +101,7 @@ public synchronized void updateCertificate(PrivateKey key, List<X509Certificate>
checkNotNull(certChain, "certChain");
privateKey = key;
this.certChain = certChain;
for (Watcher watcher : downstreamWatchers) {
for (Watcher watcher : getDownstreamWatchers()) {
sendLastCertificateUpdate(watcher);
}
}
Expand All @@ -109,39 +110,40 @@ public synchronized void updateCertificate(PrivateKey key, List<X509Certificate>
public synchronized void updateTrustedRoots(List<X509Certificate> trustedRoots) {
checkNotNull(trustedRoots, "trustedRoots");
this.trustedRoots = trustedRoots;
for (Watcher watcher : downstreamWatchers) {
for (Watcher watcher : getDownstreamWatchers()) {
sendLastTrustedRootsUpdate(watcher);
}
}

@Override
public void updateSpiffeTrustMap(Map<String, List<X509Certificate>> spiffeTrustMap) {
public synchronized void updateSpiffeTrustMap(
Map<String, List<X509Certificate>> spiffeTrustMap) {
this.spiffeTrustMap = spiffeTrustMap;
for (Watcher watcher : downstreamWatchers) {
for (Watcher watcher : getDownstreamWatchers()) {
sendLastSpiffeTrustMapUpdate(watcher);
}
}

@Override
public synchronized void onError(Status errorStatus) {
for (Watcher watcher : downstreamWatchers) {
for (Watcher watcher : getDownstreamWatchers()) {
watcher.onError(errorStatus);
}
}

X509Certificate getLastIdentityCert() {
synchronized X509Certificate getLastIdentityCert() {
if (certChain != null && !certChain.isEmpty()) {
return certChain.get(0);
}
return null;
}

void close() {
synchronized void close() {
downstreamWatchers.clear();
clearValues();
}

void clearValues() {
synchronized void clearValues() {
privateKey = null;
certChain = null;
trustedRoots = null;
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,192 @@
/*
* 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.xds.internal.security.certprovider;

import static com.google.common.truth.Truth.assertThat;
import static org.mockito.Mockito.doAnswer;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;

import com.google.common.collect.ImmutableList;
import com.google.common.collect.ImmutableMap;
import io.grpc.Status;
import io.grpc.xds.internal.security.certprovider.CertificateProvider.DistributorWatcher;
import io.grpc.xds.internal.security.certprovider.CertificateProvider.Watcher;
import java.security.PrivateKey;
import java.security.cert.X509Certificate;
import java.util.List;
import java.util.Map;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.Future;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.TimeoutException;
import java.util.concurrent.atomic.AtomicBoolean;
import org.junit.Test;
import org.junit.runner.RunWith;
import org.junit.runners.JUnit4;

/** Tests notification delivery when callbacks change certificate provider subscriptions. */
@RunWith(JUnit4.class)
public class CertificateProviderTest {
private final DistributorWatcher distributor = new DistributorWatcher();
private final Watcher watcher1 = mock(Watcher.class);
private final Watcher watcher2 = mock(Watcher.class);
private final Map<String, List<X509Certificate>> trustMap =
ImmutableMap.of("example.org", ImmutableList.of(mock(X509Certificate.class)));

@Test
public void updateSpiffeTrustMap_doesNotOverlapCertificateReplay() throws Exception {
PrivateKey key = mock(PrivateKey.class);
List<X509Certificate> certChain = ImmutableList.of(mock(X509Certificate.class));
distributor.updateCertificate(key, certChain);
CountDownLatch replayStarted = new CountDownLatch(1);
CountDownLatch finishReplay = new CountDownLatch(1);
CountDownLatch updateStarted = new CountDownLatch(1);
AtomicBoolean replayingCertificate = new AtomicBoolean();
AtomicBoolean overlappingNotifications = new AtomicBoolean();
doAnswer(invocation -> {
replayingCertificate.set(true);
replayStarted.countDown();
assertThat(finishReplay.await(10, TimeUnit.SECONDS)).isTrue();
replayingCertificate.set(false);
return null;
}).when(watcher1).updateCertificate(key, certChain);
doAnswer(invocation -> {
overlappingNotifications.compareAndSet(false, replayingCertificate.get());
return null;
}).when(watcher1).updateSpiffeTrustMap(trustMap);

ExecutorService executor = Executors.newFixedThreadPool(2);
try {
Future<?> registration = executor.submit(() -> distributor.addWatcher(watcher1));
assertThat(replayStarted.await(10, TimeUnit.SECONDS)).isTrue();
Future<?> update = executor.submit(() -> {
updateStarted.countDown();
distributor.updateSpiffeTrustMap(trustMap);
});
assertThat(updateStarted.await(10, TimeUnit.SECONDS)).isTrue();
try {
// A serialized update waits for registration to finish replaying the certificate.
update.get(1, TimeUnit.SECONDS);
} catch (TimeoutException expected) {
// Release registration below so both tasks can finish before checking delivery.
} finally {
finishReplay.countDown();
}
registration.get(10, TimeUnit.SECONDS);
update.get(10, TimeUnit.SECONDS);

assertThat(overlappingNotifications.get()).isFalse();
verify(watcher1).updateSpiffeTrustMap(trustMap);
} finally {
finishReplay.countDown();
executor.shutdownNow();
assertThat(executor.awaitTermination(10, TimeUnit.SECONDS)).isTrue();
}
}

@Test
public void updateCertificate_watcherRemovedDuringNotification() {
distributor.addWatcher(watcher1);
distributor.addWatcher(watcher2);
PrivateKey key = mock(PrivateKey.class);
List<X509Certificate> certChain = ImmutableList.of(mock(X509Certificate.class));
Watcher firstWatcher = distributor.getDownstreamWatchers().iterator().next();
doAnswer(invocation -> {
distributor.removeWatcher(firstWatcher);
return null;
}).when(firstWatcher).updateCertificate(key, certChain);

distributor.updateCertificate(key, certChain);

verify(watcher1).updateCertificate(key, certChain);
verify(watcher2).updateCertificate(key, certChain);
}

@Test
public void updateTrustedRoots_watcherRemovedDuringNotification() {
distributor.addWatcher(watcher1);
distributor.addWatcher(watcher2);
List<X509Certificate> trustedRoots = ImmutableList.of(mock(X509Certificate.class));
Watcher firstWatcher = distributor.getDownstreamWatchers().iterator().next();
doAnswer(invocation -> {
distributor.removeWatcher(firstWatcher);
return null;
}).when(firstWatcher).updateTrustedRoots(trustedRoots);

distributor.updateTrustedRoots(trustedRoots);

verify(watcher1).updateTrustedRoots(trustedRoots);
verify(watcher2).updateTrustedRoots(trustedRoots);
}

@Test
public void updateSpiffeTrustMap_watcherAddedDuringNotification() {
distributor.addWatcher(watcher1);
distributor.addWatcher(watcher2);
Watcher addedWatcher = mock(Watcher.class);
// Select the first callback without depending on HashSet's iteration order. At least one
// existing watcher still needs its initial trust map when this callback changes the set.
Watcher firstWatcher = distributor.getDownstreamWatchers().iterator().next();
doAnswer(invocation -> {
distributor.addWatcher(addedWatcher);
return null;
}).when(firstWatcher).updateSpiffeTrustMap(trustMap);

distributor.updateSpiffeTrustMap(trustMap);

verify(watcher1).updateSpiffeTrustMap(trustMap);
verify(watcher2).updateSpiffeTrustMap(trustMap);
verify(addedWatcher).updateSpiffeTrustMap(trustMap);
}

@Test
public void updateSpiffeTrustMap_watcherRemovedDuringNotification() {
distributor.addWatcher(watcher1);
distributor.addWatcher(watcher2);
Watcher firstWatcher = distributor.getDownstreamWatchers().iterator().next();
doAnswer(invocation -> {
distributor.removeWatcher(firstWatcher);
return null;
}).when(firstWatcher).updateSpiffeTrustMap(trustMap);

distributor.updateSpiffeTrustMap(trustMap);

verify(watcher1).updateSpiffeTrustMap(trustMap);
verify(watcher2).updateSpiffeTrustMap(trustMap);
assertThat(distributor.getDownstreamWatchers()).doesNotContain(firstWatcher);
}

@Test
public void onError_watcherRemovedDuringNotification() {
distributor.addWatcher(watcher1);
distributor.addWatcher(watcher2);
Watcher firstWatcher = distributor.getDownstreamWatchers().iterator().next();
doAnswer(invocation -> {
distributor.removeWatcher(firstWatcher);
return null;
}).when(firstWatcher).onError(Status.UNAVAILABLE);

distributor.onError(Status.UNAVAILABLE);

verify(watcher1).onError(Status.UNAVAILABLE);
verify(watcher2).onError(Status.UNAVAILABLE);
assertThat(distributor.getDownstreamWatchers()).doesNotContain(firstWatcher);
}
}
Loading