diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/LocationAwareSharedBackendReplicaHarnessTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/LocationAwareSharedBackendReplicaHarnessTest.java index d55e95653e21..3e0924591cb1 100644 --- a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/LocationAwareSharedBackendReplicaHarnessTest.java +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/LocationAwareSharedBackendReplicaHarnessTest.java @@ -16,6 +16,7 @@ package com.google.cloud.spanner; +import static org.awaitility.Awaitility.await; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotEquals; import static org.junit.Assert.assertTrue; @@ -49,6 +50,7 @@ import io.grpc.Status; import io.grpc.StatusRuntimeException; import io.grpc.protobuf.ProtoUtils; +import java.time.Duration; import java.util.ArrayList; import java.util.Arrays; import java.util.List; @@ -573,8 +575,13 @@ private static boolean drain(ResultSet resultSet) { return sawRow; } + private static void waitForAllReplicasConnected(SharedBackendReplicaHarness harness) { + await().atMost(Duration.ofSeconds(10)).until(harness::allReplicasConnected); + } + private static int waitForReplicaRoutedRead( DatabaseClient client, SharedBackendReplicaHarness harness) throws InterruptedException { + waitForAllReplicasConnected(harness); long deadlineNanos = System.nanoTime() + TimeUnit.SECONDS.toNanos(10); while (System.nanoTime() < deadlineNanos) { try (ResultSet resultSet = diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/SharedBackendReplicaHarness.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/SharedBackendReplicaHarness.java index 7aa5eb88c3e0..7395a55d44fd 100644 --- a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/SharedBackendReplicaHarness.java +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/SharedBackendReplicaHarness.java @@ -34,12 +34,14 @@ import com.google.spanner.v1.Session; import com.google.spanner.v1.SpannerGrpc; import com.google.spanner.v1.Transaction; +import io.grpc.Attributes; import io.grpc.Metadata; import io.grpc.Server; import io.grpc.ServerCall; import io.grpc.ServerCallHandler; import io.grpc.ServerInterceptor; import io.grpc.ServerInterceptors; +import io.grpc.ServerTransportFilter; import io.grpc.netty.shaded.io.grpc.netty.NettyServerBuilder; import io.grpc.stub.StreamObserver; import java.io.Closeable; @@ -50,6 +52,7 @@ import java.util.HashMap; import java.util.List; import java.util.Map; +import java.util.concurrent.atomic.AtomicInteger; /** Shared-backend replica harness for end-to-end location-aware routing tests. */ final class SharedBackendReplicaHarness implements Closeable { @@ -71,11 +74,24 @@ static final class HookedReplicaSpannerService extends SpannerGrpc.SpannerImplBa private final Map> methodErrors = new HashMap<>(); private final Map> requests = new HashMap<>(); private final Map> requestIds = new HashMap<>(); + private final AtomicInteger activeConnections = new AtomicInteger(); private HookedReplicaSpannerService(MockSpannerServiceImpl backend) { this.backend = backend; } + void recordConnectionReady() { + activeConnections.incrementAndGet(); + } + + void recordConnectionTerminated() { + activeConnections.decrementAndGet(); + } + + boolean hasConnected() { + return activeConnections.get() > 0; + } + synchronized void putMethodErrors(String method, Throwable... errors) { ArrayDeque queue = new ArrayDeque<>(); for (Throwable error : errors) { @@ -278,12 +294,35 @@ public ServerCall.Listener interceptCall( Server server = NettyServerBuilder.forAddress(address) .addService(ServerInterceptors.intercept(service, interceptor)) + .addTransportFilter( + new ServerTransportFilter() { + @Override + public Attributes transportReady(Attributes transportAttrs) { + service.recordConnectionReady(); + return super.transportReady(transportAttrs); + } + + @Override + public void transportTerminated(Attributes transportAttrs) { + service.recordConnectionTerminated(); + super.transportTerminated(transportAttrs); + } + }) .build() .start(); servers.add(server); return "localhost:" + server.getPort(); } + boolean allReplicasConnected() { + for (HookedReplicaSpannerService replica : replicas) { + if (!replica.hasConnected()) { + return false; + } + } + return true; + } + void clearRequests() { defaultReplica.clearRequests(); for (HookedReplicaSpannerService replica : replicas) {