diff --git a/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/GapicSpannerRpc.java b/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/GapicSpannerRpc.java index 807eac6f11b..608d9e23d53 100644 --- a/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/GapicSpannerRpc.java +++ b/google-cloud-spanner/src/main/java/com/google/cloud/spanner/spi/v1/GapicSpannerRpc.java @@ -77,7 +77,6 @@ import com.google.common.base.MoreObjects; import com.google.common.base.Preconditions; import com.google.common.collect.ImmutableList; -import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; import com.google.common.util.concurrent.RateLimiter; import com.google.common.util.concurrent.ThreadFactoryBuilder; @@ -161,6 +160,7 @@ import java.net.URLDecoder; import java.nio.charset.StandardCharsets; import java.util.Comparator; +import java.util.HashMap; import java.util.LinkedList; import java.util.List; import java.util.Map; @@ -244,6 +244,8 @@ private void awaitTermination() throws InterruptedException { private static final int GRPC_KEEPALIVE_SECONDS = 2 * 60; private static final String USER_AGENT_KEY = "user-agent"; private static final String CLIENT_LIBRARY_LANGUAGE = "spanner-java"; + public static final String DEFAULT_USER_AGENT = + CLIENT_LIBRARY_LANGUAGE + "/" + GaxProperties.getLibraryVersion(GapicSpannerRpc.class); private final ManagedInstantiatingExecutorProvider executorProvider; private boolean rpcIsClosed; @@ -305,18 +307,11 @@ public GapicSpannerRpc(final SpannerOptions options) { GaxGrpcProperties.getGrpcTokenName(), GaxGrpcProperties.getGrpcVersion()) .build(); - HeaderProvider mergedHeaderProvider = options.getMergedHeaderProvider(internalHeaderProvider); - Map headersWithUserAgent = - ImmutableMap.builder() - .put( - USER_AGENT_KEY, - CLIENT_LIBRARY_LANGUAGE - + "/" - + GaxProperties.getLibraryVersion(GapicSpannerRpc.class)) - .putAll(mergedHeaderProvider.getHeaders()) - .build(); + final HeaderProvider mergedHeaderProvider = + options.getMergedHeaderProvider(internalHeaderProvider); final HeaderProvider headerProviderWithUserAgent = - FixedHeaderProvider.create(headersWithUserAgent); + headerProviderWithUserAgentFrom(mergedHeaderProvider); + this.metadataProvider = SpannerMetadataProvider.create( headerProviderWithUserAgent.getHeaders(), @@ -494,6 +489,16 @@ public UnaryCallable createUnaryCalla } } + private static HeaderProvider headerProviderWithUserAgentFrom(HeaderProvider headerProvider) { + final Map headersWithUserAgent = new HashMap<>(headerProvider.getHeaders()); + final String userAgent = headersWithUserAgent.get(USER_AGENT_KEY); + headersWithUserAgent.put( + USER_AGENT_KEY, + userAgent == null ? DEFAULT_USER_AGENT : userAgent + " " + DEFAULT_USER_AGENT); + + return FixedHeaderProvider.create(headersWithUserAgent); + } + private static void checkEmulatorConnection( SpannerOptions options, TransportChannelProvider channelProvider, diff --git a/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/GapicSpannerRpcTest.java b/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/GapicSpannerRpcTest.java index 84aaa91bcf5..5a9dec72edc 100644 --- a/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/GapicSpannerRpcTest.java +++ b/google-cloud-spanner/src/test/java/com/google/cloud/spanner/spi/v1/GapicSpannerRpcTest.java @@ -24,7 +24,9 @@ import static org.junit.Assume.assumeTrue; import com.google.api.core.ApiFunction; +import com.google.api.gax.core.GaxProperties; import com.google.api.gax.rpc.ApiCallContext; +import com.google.api.gax.rpc.HeaderProvider; import com.google.auth.oauth2.AccessToken; import com.google.auth.oauth2.OAuth2Credentials; import com.google.cloud.spanner.DatabaseAdminClient; @@ -151,6 +153,8 @@ public class GapicSpannerRpcTest { private Server server; private InetSocketAddress address; private final Map optionsMap = new HashMap<>(); + private Metadata seenHeaders; + private String defaultUserAgent; @BeforeClass public static void checkNotEmulator() { @@ -161,6 +165,7 @@ public static void checkNotEmulator() { @Before public void startServer() throws IOException { + defaultUserAgent = "spanner-java/" + GaxProperties.getLibraryVersion(GapicSpannerRpc.class); mockSpanner = new MockSpannerServiceImpl(); mockSpanner.setAbortProbability(0.0D); // We don't want any unpredictable aborted transactions. mockSpanner.putStatementResult(StatementResult.query(SELECT1AND2, SELECT1_RESULTSET)); @@ -183,6 +188,7 @@ public ServerCall.Listener interceptCall( ServerCall call, Metadata headers, ServerCallHandler next) { + seenHeaders = headers; String auth = headers.get(Key.of("authorization", Metadata.ASCII_STRING_MARSHALLER)); assertThat(auth).isEqualTo("Bearer " + VARIABLE_OAUTH_TOKEN); @@ -502,6 +508,46 @@ public void testAdminRequestsLimitExceededRetryAlgorithm() { assertThat(alg.shouldRetry(new Exception("random exception"), null)).isFalse(); } + @Test + public void testDefaultUserAgent() { + final SpannerOptions options = createSpannerOptions(); + final Spanner spanner = options.getService(); + final DatabaseClient databaseClient = + spanner.getDatabaseClient(DatabaseId.of("[PROJECT]", "[INSTANCE]", "[DATABASE]")); + + try (final ResultSet rs = databaseClient.singleUse().executeQuery(SELECT1AND2)) { + rs.next(); + } + + assertThat(seenHeaders.get(Key.of("user-agent", Metadata.ASCII_STRING_MARSHALLER))) + .contains(defaultUserAgent); + } + + @Test + public void testCustomUserAgent() { + final HeaderProvider userAgentHeaderProvider = + new HeaderProvider() { + @Override + public Map getHeaders() { + final Map headers = new HashMap<>(); + headers.put("user-agent", "test-agent"); + return headers; + } + }; + final SpannerOptions options = + createSpannerOptions().toBuilder().setHeaderProvider(userAgentHeaderProvider).build(); + final Spanner spanner = options.getService(); + final DatabaseClient databaseClient = + spanner.getDatabaseClient(DatabaseId.of("[PROJECT]", "[INSTANCE]", "[DATABASE]")); + + try (final ResultSet rs = databaseClient.singleUse().executeQuery(SELECT1AND2)) { + rs.next(); + } + + assertThat(seenHeaders.get(Key.of("user-agent", Metadata.ASCII_STRING_MARSHALLER))) + .contains("test-agent " + defaultUserAgent); + } + @SuppressWarnings("rawtypes") private SpannerOptions createSpannerOptions() { String endpoint = address.getHostString() + ":" + server.getPort();