diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/HeaderInterceptor.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/HeaderInterceptor.java index 9374334a86ad..6f69a1ea97fd 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/HeaderInterceptor.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/HeaderInterceptor.java @@ -25,13 +25,23 @@ import static com.google.cloud.spanner.spi.v1.SpannerRpcViews.SPANNER_GFE_LATENCY; import com.google.api.gax.tracing.ApiTracer; -import com.google.cloud.spanner.*; +import com.google.cloud.spanner.BuiltInMetricsConstant; +import com.google.cloud.spanner.CompositeTracer; +import com.google.cloud.spanner.SpannerExceptionFactory; +import com.google.cloud.spanner.SpannerRpcMetrics; +import com.google.cloud.spanner.XGoogSpannerRequestId; import com.google.common.cache.Cache; import com.google.common.cache.CacheBuilder; import com.google.spanner.admin.database.v1.DatabaseName; -import io.grpc.*; +import io.grpc.CallOptions; +import io.grpc.Channel; +import io.grpc.ClientCall; +import io.grpc.ClientInterceptor; import io.grpc.ForwardingClientCall.SimpleForwardingClientCall; import io.grpc.ForwardingClientCallListener.SimpleForwardingClientCallListener; +import io.grpc.Metadata; +import io.grpc.MethodDescriptor; +import io.grpc.Status; import io.grpc.alts.AltsContextUtil; import io.opencensus.stats.MeasureMap; import io.opencensus.stats.Stats; @@ -52,6 +62,7 @@ import java.util.logging.Logger; import java.util.regex.Matcher; import java.util.regex.Pattern; +import javax.annotation.Nullable; /** * Intercepts all gRPC calls to extract server-timing header. Captures GFE Latency and GFE Header @@ -89,8 +100,6 @@ class HeaderInterceptor implements ClientInterceptor { private static final Logger LOGGER = Logger.getLogger(HeaderInterceptor.class.getName()); private static final Level LEVEL = Level.INFO; private final SpannerRpcMetrics spannerRpcMetrics; - private Float gfeLatency; - private Float afeLatency; HeaderInterceptor(SpannerRpcMetrics spannerRpcMetrics) { this.spannerRpcMetrics = spannerRpcMetrics; @@ -118,48 +127,65 @@ public void start(Listener responseListener, Metadata headers) { super.start( new SimpleForwardingClientCallListener(responseListener) { + private Float gfeLatency; + private Float afeLatency; + @Override public void onHeaders(Metadata metadata) { - recordFirstResponseLatency(requestId, startedAtNanos, firstResponseRecorded); - String serverTiming = metadata.get(SERVER_TIMING_HEADER_KEY); try { - // Get gfe and afe Latency value - Map serverTimingMetrics = parseServerTimingHeader(serverTiming); - gfeLatency = serverTimingMetrics.get(GFE_TIMING_HEADER); - afeLatency = serverTimingMetrics.get(AFE_TIMING_HEADER); - } catch (NumberFormatException e) { - LOGGER.log(LEVEL, "Invalid server-timing object in header: {}", serverTiming); + recordFirstResponseLatency(requestId, startedAtNanos, firstResponseRecorded); + String serverTiming = metadata.get(SERVER_TIMING_HEADER_KEY); + try { + // Get gfe and afe Latency value + Map serverTimingMetrics = + parseServerTimingHeader(serverTiming); + gfeLatency = serverTimingMetrics.get(GFE_TIMING_HEADER); + afeLatency = serverTimingMetrics.get(AFE_TIMING_HEADER); + } catch (NumberFormatException e) { + LOGGER.log( + LEVEL, "Invalid server-timing object in header: {0}", serverTiming); + } + } catch (Throwable throwable) { + LOGGER.log( + Level.WARNING, "Error processing headers in HeaderInterceptor", throwable); + } finally { + super.onHeaders(metadata); } - - super.onHeaders(metadata); } @Override public void onClose(Status status, Metadata trailers) { - // Record Built-in Metrics - boolean isDirectPathUsed = AltsContextUtil.check(getAttributes()); - boolean isAfeEnabled = GapicSpannerRpc.isEnableAFEServerTiming(); - recordSpan(span, requestId); - recordCustomMetrics(tagContext, attributes, isDirectPathUsed); - Map builtInMetricsAttributes = new HashMap<>(); try { - builtInMetricsAttributes = - new HashMap<>(getBuiltInMetricAttributes(key, databaseName)); - } catch (ExecutionException e) { - LOGGER.log( - LEVEL, "Unable to get built-in metric attributes {}", e.getMessage()); + // Record Built-in Metrics + boolean isDirectPathUsed = AltsContextUtil.check(getAttributes()); + boolean isAfeEnabled = GapicSpannerRpc.isEnableAFEServerTiming(); + recordSpan(span, requestId, gfeLatency, afeLatency); + recordCustomMetrics(tagContext, attributes, isDirectPathUsed, gfeLatency); + Map builtInMetricsAttributes = new HashMap<>(); + try { + builtInMetricsAttributes = + new HashMap<>(getBuiltInMetricAttributes(key, databaseName)); + } catch (ExecutionException e) { + LOGGER.log( + LEVEL, "Unable to get built-in metric attributes {0}", e.getMessage()); + } + if (status.isOk()) { + recordFirstResponseLatency(requestId, startedAtNanos, firstResponseRecorded); + } + recordBuiltInMetrics( + compositeTracer, + builtInMetricsAttributes, + requestId, + isDirectPathUsed, + isAfeEnabled, + gfeLatency, + afeLatency); + } catch (Throwable throwable) { + LOGGER.log(Level.WARNING, "Error recording metrics in onClose", throwable); + } finally { + RequestIdTargetTracker.remove(requestId); + super.onClose(status, trailers); } - if (status.isOk()) { - recordFirstResponseLatency(requestId, startedAtNanos, firstResponseRecorded); - } - recordBuiltInMetrics( - compositeTracer, - builtInMetricsAttributes, - requestId, - isDirectPathUsed, - isAfeEnabled); - RequestIdTargetTracker.remove(requestId); - super.onClose(status, trailers); } }, headers); @@ -172,11 +198,14 @@ public void onClose(Status status, Metadata trailers) { } private void recordCustomMetrics( - TagContext tagContext, Attributes attributes, Boolean isDirectPathUsed) { + TagContext tagContext, + Attributes attributes, + Boolean isDirectPathUsed, + @Nullable Float gfeLatency) { // Record OpenCensus and Custom OpenTelemetry Metrics MeasureMap measureMap = STATS_RECORDER.newMeasureMap(); - if (!isDirectPathUsed) { + if (!Boolean.TRUE.equals(isDirectPathUsed)) { if (gfeLatency != null) { long gfeVal = gfeLatency.longValue(); measureMap.put(SPANNER_GFE_LATENCY, gfeVal); @@ -191,7 +220,11 @@ private void recordCustomMetrics( measureMap.record(tagContext); } - private void recordSpan(Span span, String requestId) { + private void recordSpan( + @Nullable Span span, + @Nullable String requestId, + @Nullable Float gfeLatency, + @Nullable Float afeLatency) { if (span != null) { if (gfeLatency != null) { span.setAttribute("gfe_latency", gfeLatency.toString()); @@ -199,23 +232,33 @@ private void recordSpan(Span span, String requestId) { if (afeLatency != null) { span.setAttribute("afe_latency", afeLatency.toString()); } - span.setAttribute(XGoogSpannerRequestId.REQUEST_ID_HEADER_NAME, requestId); + if (requestId != null) { + span.setAttribute(XGoogSpannerRequestId.REQUEST_ID_HEADER_NAME, requestId); + } } } private void recordBuiltInMetrics( - CompositeTracer compositeTracer, + @Nullable CompositeTracer compositeTracer, Map builtInMetricsAttributes, - String requestId, + @Nullable String requestId, Boolean isDirectPathUsed, - Boolean isAfeEnabled) { + Boolean isAfeEnabled, + @Nullable Float gfeLatency, + @Nullable Float afeLatency) { if (compositeTracer != null) { - builtInMetricsAttributes.put(BuiltInMetricsConstant.REQUEST_ID_KEY.getKey(), requestId); + if (requestId != null) { + builtInMetricsAttributes.put(BuiltInMetricsConstant.REQUEST_ID_KEY.getKey(), requestId); + } builtInMetricsAttributes.put( - BuiltInMetricsConstant.DIRECT_PATH_USED_KEY.getKey(), Boolean.toString(isDirectPathUsed)); + BuiltInMetricsConstant.DIRECT_PATH_USED_KEY.getKey(), + Boolean.toString(Boolean.TRUE.equals(isDirectPathUsed))); compositeTracer.addAttributes(builtInMetricsAttributes); compositeTracer.recordServerTimingHeaderMetrics( - gfeLatency, afeLatency, isDirectPathUsed, isAfeEnabled); + gfeLatency, + afeLatency, + Boolean.TRUE.equals(isDirectPathUsed), + Boolean.TRUE.equals(isAfeEnabled)); } } diff --git a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/SpannerErrorInterceptor.java b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/SpannerErrorInterceptor.java index 9c3b2af2b069..9535e751b0fa 100644 --- a/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/SpannerErrorInterceptor.java +++ b/java-spanner/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/SpannerErrorInterceptor.java @@ -118,10 +118,13 @@ public void onClose(Status status, Metadata trailers) { if (trailers.containsKey(RETRY_INFO_KEY)) { status = status.augmentDescription(trailers.get(RETRY_INFO_KEY).toString()); } - } catch (IllegalArgumentException e) { + } catch (Throwable throwable) { // Messages could be invalid if, say, some invalid UTF8 is echoed back in some - // error text. - logger.log(Level.WARNING, "Invalid protocol message in metadata", e); + // error text, or if an unexpected exception occurs during metadata inspection. + logger.log( + Level.WARNING, + "Error processing error details in SpannerErrorInterceptor", + throwable); } finally { super.onClose(status, trailers); } diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/HeaderInterceptorTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/HeaderInterceptorTest.java new file mode 100644 index 000000000000..ba4d9fcf5106 --- /dev/null +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/HeaderInterceptorTest.java @@ -0,0 +1,345 @@ +/* + * Copyright 2026 Google LLC + * + * 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 com.google.cloud.spanner.spi.v1; + +import static com.google.api.gax.grpc.GrpcCallContext.TRACER_KEY; +import static com.google.cloud.spanner.XGoogSpannerRequestId.REQUEST_ID_HEADER_KEY; +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNotNull; +import static org.junit.Assert.assertNull; +import static org.junit.Assert.assertTrue; + +import com.google.cloud.spanner.CompositeTracer; +import com.google.cloud.spanner.SpannerRpcMetrics; +import com.google.common.collect.ImmutableList; +import io.grpc.CallOptions; +import io.grpc.Channel; +import io.grpc.ClientCall; +import io.grpc.Metadata; +import io.grpc.MethodDescriptor; +import io.grpc.MethodDescriptor.Marshaller; +import io.grpc.Status; +import io.opentelemetry.api.OpenTelemetry; +import java.io.InputStream; +import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.After; +import org.junit.Before; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +/** Unit tests for {@link HeaderInterceptor}. */ +@RunWith(JUnit4.class) +public class HeaderInterceptorTest { + + private static final Metadata.Key SERVER_TIMING_HEADER_KEY = + Metadata.Key.of("server-timing", Metadata.ASCII_STRING_MARSHALLER); + private static final Metadata.Key GOOGLE_CLOUD_RESOURCE_PREFIX_KEY = + Metadata.Key.of("google-cloud-resource-prefix", Metadata.ASCII_STRING_MARSHALLER); + + private HeaderInterceptor interceptor; + + @Before + public void setUp() { + interceptor = new HeaderInterceptor(new SpannerRpcMetrics(OpenTelemetry.noop())); + RequestIdTargetTracker.clear(); + } + + @After + public void tearDown() { + RequestIdTargetTracker.clear(); + } + + @Test + public void testPerCallLatencyIsolation() { + TestCompositeTracer tracer1 = new TestCompositeTracer(); + TestCompositeTracer tracer2 = new TestCompositeTracer(); + + CallOptions callOptions1 = CallOptions.DEFAULT.withOption(TRACER_KEY, tracer1); + CallOptions callOptions2 = CallOptions.DEFAULT.withOption(TRACER_KEY, tracer2); + + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + + ClientCall call1 = + interceptor.interceptCall(methodDescriptor, callOptions1, channel); + ClientCall call2 = + interceptor.interceptCall(methodDescriptor, callOptions2, channel); + + Metadata headers1 = createDefaultHeaders("1.0000000000000001.1.1.1.1"); + Metadata headers2 = createDefaultHeaders("1.0000000000000002.1.1.2.1"); + + CapturingListener responseListener1 = new CapturingListener<>(); + CapturingListener responseListener2 = new CapturingListener<>(); + + call1.start(responseListener1, headers1); + call2.start(responseListener2, headers2); + + Metadata responseHeaders1 = new Metadata(); + responseHeaders1.put(SERVER_TIMING_HEADER_KEY, "gfet4t7; dur=123.45, afe; dur=67.89"); + + Metadata responseHeaders2 = new Metadata(); + responseHeaders2.put(SERVER_TIMING_HEADER_KEY, "gfet4t7; dur=999.0, afe; dur=888.0"); + + // Interleave headers: call1 then call2 + channel.lastListener.onHeaders(responseHeaders2); // Call 2 listener + channel.capturedListeners.get(0).onHeaders(responseHeaders1); // Call 1 listener + + // Close call 2 first, then call 1 + channel.lastListener.onClose(Status.OK, new Metadata()); + channel.capturedListeners.get(0).onClose(Status.OK, new Metadata()); + + // Verify Call 1 received its own latencies, not Call 2's + assertEquals(Float.valueOf(123.45f), tracer1.recordedGfeLatency); + assertEquals(Float.valueOf(67.89f), tracer1.recordedAfeLatency); + + // Verify Call 2 received its own latencies + assertEquals(Float.valueOf(999.0f), tracer2.recordedGfeLatency); + assertEquals(Float.valueOf(888.0f), tracer2.recordedAfeLatency); + + assertTrue(responseListener1.closed.get()); + assertTrue(responseListener2.closed.get()); + } + + @Test + public void testLatencyNotPollutedWhenSecondCallHasNoHeader() { + TestCompositeTracer tracer1 = new TestCompositeTracer(); + TestCompositeTracer tracer2 = new TestCompositeTracer(); + + CallOptions callOptions1 = CallOptions.DEFAULT.withOption(TRACER_KEY, tracer1); + CallOptions callOptions2 = CallOptions.DEFAULT.withOption(TRACER_KEY, tracer2); + + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + + ClientCall call1 = + interceptor.interceptCall(methodDescriptor, callOptions1, channel); + ClientCall call2 = + interceptor.interceptCall(methodDescriptor, callOptions2, channel); + + Metadata headers1 = createDefaultHeaders("1.0000000000000001.1.1.1.1"); + Metadata headers2 = createDefaultHeaders("1.0000000000000002.1.1.2.1"); + + CapturingListener responseListener1 = new CapturingListener<>(); + CapturingListener responseListener2 = new CapturingListener<>(); + + call1.start(responseListener1, headers1); + call2.start(responseListener2, headers2); + + Metadata responseHeaders1 = new Metadata(); + responseHeaders1.put(SERVER_TIMING_HEADER_KEY, "gfet4t7; dur=150.0"); + + Metadata responseHeaders2 = new Metadata(); // No server-timing header + + channel.capturedListeners.get(0).onHeaders(responseHeaders1); + channel.capturedListeners.get(1).onHeaders(responseHeaders2); + + channel.capturedListeners.get(1).onClose(Status.OK, new Metadata()); + channel.capturedListeners.get(0).onClose(Status.OK, new Metadata()); + + assertEquals(Float.valueOf(150.0f), tracer1.recordedGfeLatency); + assertNull(tracer2.recordedGfeLatency); + assertNull(tracer2.recordedAfeLatency); + } + + @Test + public void testOnCloseGuaranteesDownstreamNotificationEvenOnException() { + TestCompositeTracer throwingTracer = + new TestCompositeTracer() { + @Override + public void recordServerTimingHeaderMetrics( + Float gfeLatency, Float afeLatency, boolean isDirectPathUsed, boolean isAfeEnabled) { + throw new RuntimeException("Simulated metric recording failure"); + } + }; + + CallOptions callOptions = CallOptions.DEFAULT.withOption(TRACER_KEY, throwingTracer); + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + + String requestId = "1.0000000000000001.1.1.1.1"; + RequestIdTargetTracker.record(requestId, "test-database", "endpoint-1", 100L, false); + assertNotNull(RequestIdTargetTracker.get(requestId)); + + ClientCall call = + interceptor.interceptCall(methodDescriptor, callOptions, channel); + CapturingListener responseListener = new CapturingListener<>(); + call.start(responseListener, createDefaultHeaders(requestId)); + + // Deliver onClose - even though metric recording throws, onClose must propagate downstream + channel.lastListener.onClose(Status.OK, new Metadata()); + + assertTrue( + "Downstream listener must receive onClose despite metric recording error", + responseListener.closed.get()); + assertEquals(Status.OK, responseListener.closedStatus.get()); + + // RequestIdTargetTracker must still be cleaned up in finally block + assertNull( + "RequestIdTargetTracker must clean up tracking entry in finally", + RequestIdTargetTracker.get(requestId)); + } + + @Test + public void testOnHeadersGuaranteesDownstreamNotificationEvenOnException() { + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + + ClientCall call = + interceptor.interceptCall(methodDescriptor, CallOptions.DEFAULT, channel); + CapturingListener responseListener = new CapturingListener<>(); + call.start(responseListener, createDefaultHeaders("1.0000000000000001.1.1.1.1")); + + Metadata corruptedHeaders = new Metadata(); + corruptedHeaders.put(SERVER_TIMING_HEADER_KEY, "gfet4t7; dur=invalid_number"); + + channel.lastListener.onHeaders(corruptedHeaders); + + assertTrue( + "Downstream listener must receive onHeaders despite parsing anomaly", + responseListener.headersReceived.get()); + } + + @Test + public void testCallWithoutRequestIdDoesNotThrowException() { + TestCompositeTracer tracer = new TestCompositeTracer(); + CallOptions callOptions = CallOptions.DEFAULT.withOption(TRACER_KEY, tracer); + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + + ClientCall call = + interceptor.interceptCall(methodDescriptor, callOptions, channel); + CapturingListener responseListener = new CapturingListener<>(); + + // Start with headers that do NOT contain REQUEST_ID_HEADER_KEY + Metadata headersWithoutRequestId = new Metadata(); + headersWithoutRequestId.put( + GOOGLE_CLOUD_RESOURCE_PREFIX_KEY, + "projects/test-project/instances/test-instance/databases/test-database"); + call.start(responseListener, headersWithoutRequestId); + + Metadata responseHeaders = new Metadata(); + responseHeaders.put(SERVER_TIMING_HEADER_KEY, "gfet4t7; dur=45.0, afe; dur=20.0"); + channel.lastListener.onHeaders(responseHeaders); + channel.lastListener.onClose(Status.OK, new Metadata()); + + assertEquals(Float.valueOf(45.0f), tracer.recordedGfeLatency); + assertEquals(Float.valueOf(20.0f), tracer.recordedAfeLatency); + assertTrue(responseListener.closed.get()); + assertEquals(Status.OK, responseListener.closedStatus.get()); + } + + private static Metadata createDefaultHeaders(String requestId) { + Metadata headers = new Metadata(); + headers.put(REQUEST_ID_HEADER_KEY, requestId); + headers.put( + GOOGLE_CLOUD_RESOURCE_PREFIX_KEY, + "projects/test-project/instances/test-instance/databases/test-database"); + return headers; + } + + private static MethodDescriptor createMethodDescriptor() { + return MethodDescriptor.newBuilder() + .setType(MethodDescriptor.MethodType.UNARY) + .setFullMethodName("google.spanner.v1.Spanner/ExecuteSql") + .setRequestMarshaller(new FakeMarshaller<>()) + .setResponseMarshaller(new FakeMarshaller<>()) + .build(); + } + + private static class TestCompositeTracer extends CompositeTracer { + Float recordedGfeLatency; + Float recordedAfeLatency; + + TestCompositeTracer() { + super(ImmutableList.of()); + } + + @Override + public void recordServerTimingHeaderMetrics( + Float gfeLatency, Float afeLatency, boolean isDirectPathUsed, boolean isAfeEnabled) { + this.recordedGfeLatency = gfeLatency; + this.recordedAfeLatency = afeLatency; + } + } + + private static class FakeMarshaller implements Marshaller { + @Override + public InputStream stream(T value) { + return null; + } + + @Override + public T parse(InputStream stream) { + return null; + } + } + + private static class FakeChannel extends Channel { + final List> capturedListeners = new CopyOnWriteArrayList<>(); + ClientCall.Listener lastListener; + + @Override + public String authority() { + return "fake-authority"; + } + + @Override + public ClientCall newCall( + MethodDescriptor methodDescriptor, CallOptions callOptions) { + return new ClientCall() { + @Override + public void start(Listener responseListener, Metadata headers) { + capturedListeners.add(responseListener); + lastListener = responseListener; + } + + @Override + public void request(int numMessages) {} + + @Override + public void cancel(String message, Throwable cause) {} + + @Override + public void halfClose() {} + + @Override + public void sendMessage(ReqT message) {} + }; + } + } + + private static class CapturingListener extends ClientCall.Listener { + final AtomicBoolean headersReceived = new AtomicBoolean(false); + final AtomicBoolean closed = new AtomicBoolean(false); + final AtomicReference closedStatus = new AtomicReference<>(); + + @Override + public void onHeaders(Metadata headers) { + headersReceived.set(true); + } + + @Override + public void onClose(Status status, Metadata trailers) { + closed.set(true); + closedStatus.set(status); + } + } +} diff --git a/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/SpannerErrorInterceptorTest.java b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/SpannerErrorInterceptorTest.java new file mode 100644 index 000000000000..91c7cf3cd057 --- /dev/null +++ b/java-spanner/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/SpannerErrorInterceptorTest.java @@ -0,0 +1,262 @@ +/* + * Copyright 2026 Google LLC + * + * 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 com.google.cloud.spanner.spi.v1; + +import static com.google.cloud.spanner.XGoogSpannerRequestId.REQUEST_ID_HEADER_KEY; +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNotNull; +import static org.junit.Assert.assertTrue; + +import com.google.rpc.LocalizedMessage; +import com.google.rpc.ResourceInfo; +import com.google.rpc.RetryInfo; +import io.grpc.CallOptions; +import io.grpc.Channel; +import io.grpc.ClientCall; +import io.grpc.Metadata; +import io.grpc.MethodDescriptor; +import io.grpc.MethodDescriptor.Marshaller; +import io.grpc.Status; +import io.grpc.protobuf.ProtoUtils; +import java.io.InputStream; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.Before; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +/** Unit tests for {@link SpannerErrorInterceptor}. */ +@RunWith(JUnit4.class) +public class SpannerErrorInterceptorTest { + + private static final Metadata.Key LOCALIZED_MESSAGE_KEY = + ProtoUtils.keyForProto(LocalizedMessage.getDefaultInstance()); + private static final Metadata.Key RESOURCE_INFO_KEY = + ProtoUtils.keyForProto(ResourceInfo.getDefaultInstance()); + private static final Metadata.Key RETRY_INFO_KEY = + ProtoUtils.keyForProto(RetryInfo.getDefaultInstance()); + + private SpannerErrorInterceptor interceptor; + + @Before + public void setUp() { + interceptor = new SpannerErrorInterceptor(); + } + + @Test + public void testOkStatusPassesThroughUnmodified() { + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + ClientCall call = + interceptor.interceptCall(methodDescriptor, CallOptions.DEFAULT, channel); + + CapturingListener responseListener = new CapturingListener<>(); + call.start(responseListener, new Metadata()); + + Metadata trailers = new Metadata(); + channel.lastListener.onClose(Status.OK, trailers); + + assertTrue(responseListener.closed.get()); + assertEquals(Status.Code.OK, responseListener.closedStatus.get().getCode()); + } + + @Test + public void testRetryableInternalErrorTranslatedToUnavailable() { + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + ClientCall call = + interceptor.interceptCall(methodDescriptor, CallOptions.DEFAULT, channel); + + CapturingListener responseListener = new CapturingListener<>(); + call.start(responseListener, new Metadata()); + + Status internalError = + Status.INTERNAL.withDescription( + "HTTP/2 error code: INTERNAL_ERROR\nReceived unexpected EOS on Data frame"); + channel.lastListener.onClose(internalError, new Metadata()); + + assertTrue(responseListener.closed.get()); + assertEquals(Status.Code.UNAVAILABLE, responseListener.closedStatus.get().getCode()); + assertTrue( + responseListener + .closedStatus + .get() + .getDescription() + .contains("Received unexpected EOS on Data frame")); + } + + @Test + public void testTrailersAugmentStatusDescription() { + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + ClientCall call = + interceptor.interceptCall(methodDescriptor, CallOptions.DEFAULT, channel); + + CapturingListener responseListener = new CapturingListener<>(); + call.start(responseListener, new Metadata()); + + Metadata trailers = new Metadata(); + LocalizedMessage localizedMessage = + LocalizedMessage.newBuilder() + .setLocale("en-US") + .setMessage("Detailed user facing error message") + .build(); + trailers.put(LOCALIZED_MESSAGE_KEY, localizedMessage); + + ResourceInfo resourceInfo = + ResourceInfo.newBuilder() + .setResourceType("type.googleapis.com/google.spanner.v1.Database") + .setResourceName("projects/p/instances/i/databases/d") + .setDescription("Database resource description") + .build(); + trailers.put(RESOURCE_INFO_KEY, resourceInfo); + + RetryInfo retryInfo = + RetryInfo.newBuilder() + .setRetryDelay(com.google.protobuf.Duration.newBuilder().setSeconds(5).build()) + .build(); + trailers.put(RETRY_INFO_KEY, retryInfo); + + Status errorStatus = Status.NOT_FOUND.withDescription("Original error"); + channel.lastListener.onClose(errorStatus, trailers); + + assertTrue(responseListener.closed.get()); + assertEquals(Status.Code.NOT_FOUND, responseListener.closedStatus.get().getCode()); + String description = responseListener.closedStatus.get().getDescription(); + assertNotNull(description); + assertTrue(description.contains("Detailed user facing error message")); + assertTrue(description.contains("Database resource description")); + assertTrue(description.contains("retry_delay")); + } + + @Test + public void testRequestIdCopiedToTrailersIfPresentInHeaders() { + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + ClientCall call = + interceptor.interceptCall(methodDescriptor, CallOptions.DEFAULT, channel); + + CapturingListener responseListener = new CapturingListener<>(); + Metadata headers = new Metadata(); + String requestId = "1.0000000000000001.1.1.1.1"; + headers.put(REQUEST_ID_HEADER_KEY, requestId); + call.start(responseListener, headers); + + Metadata trailers = new Metadata(); + channel.lastListener.onClose(Status.DEADLINE_EXCEEDED, trailers); + + assertTrue(responseListener.closed.get()); + assertEquals(requestId, responseListener.closedTrailers.get().get(REQUEST_ID_HEADER_KEY)); + } + + @Test + public void testOnCloseGuaranteesDownstreamNotificationEvenOnUnexpectedException() { + MethodDescriptor methodDescriptor = createMethodDescriptor(); + FakeChannel channel = new FakeChannel(); + ClientCall call = + interceptor.interceptCall(methodDescriptor, CallOptions.DEFAULT, channel); + + CapturingListener responseListener = new CapturingListener<>(); + Metadata headers = new Metadata(); + headers.put(REQUEST_ID_HEADER_KEY, "1.0000000000000001.1.1.1.1"); + call.start(responseListener, headers); + + // Create trailers with invalid binary protobuf payload for LocalizedMessage + Metadata trailers = new Metadata(); + Metadata.Key rawBinaryKey = + Metadata.Key.of("google.rpc.localizedmessage-bin", Metadata.BINARY_BYTE_MARSHALLER); + trailers.put(rawBinaryKey, new byte[] {(byte) 0xFF, (byte) 0xFF, (byte) 0xFF}); + + // Deliver onClose - must not throw and must reach downstream listener + channel.lastListener.onClose(Status.DATA_LOSS.withDescription("Data loss error"), trailers); + + assertTrue( + "Downstream listener must receive onClose despite unexpected exception in interceptor", + responseListener.closed.get()); + assertEquals(Status.Code.DATA_LOSS, responseListener.closedStatus.get().getCode()); + } + + private static MethodDescriptor createMethodDescriptor() { + return MethodDescriptor.newBuilder() + .setType(MethodDescriptor.MethodType.UNARY) + .setFullMethodName("google.spanner.v1.Spanner/ExecuteSql") + .setRequestMarshaller(new FakeMarshaller<>()) + .setResponseMarshaller(new FakeMarshaller<>()) + .build(); + } + + private static class FakeMarshaller implements Marshaller { + @Override + public InputStream stream(T value) { + return null; + } + + @Override + public T parse(InputStream stream) { + return null; + } + } + + private static class FakeChannel extends Channel { + ClientCall.Listener lastListener; + + @Override + public String authority() { + return "fake-authority"; + } + + @Override + public ClientCall newCall( + MethodDescriptor methodDescriptor, CallOptions callOptions) { + return new ClientCall() { + @Override + public void start(Listener responseListener, Metadata headers) { + lastListener = responseListener; + } + + @Override + public void request(int numMessages) {} + + @Override + public void cancel(String message, Throwable cause) {} + + @Override + public void halfClose() {} + + @Override + public void sendMessage(ReqT message) {} + }; + } + } + + private static class CapturingListener extends ClientCall.Listener { + final AtomicBoolean closed = new AtomicBoolean(false); + final AtomicReference closedStatus = new AtomicReference<>(); + final AtomicReference closedTrailers = new AtomicReference<>(); + + @Override + public void onHeaders(Metadata headers) {} + + @Override + public void onClose(Status status, Metadata trailers) { + closed.set(true); + closedStatus.set(status); + closedTrailers.set(trailers); + } + } +}