Move named channel back into autoconfigure module

This commit moves the channel config ("NamedChannel") back to
autoconfigure module as the effort to move to core turns out to
be more trouble than it is worth at this time.

Fixes #92
This commit is contained in:
Chris Bono
2025-01-13 00:27:48 -06:00
committed by Dave Syer
parent 5ef993fe19
commit 87e26c378e
11 changed files with 550 additions and 655 deletions

View File

@@ -1,348 +0,0 @@
/*
* Copyright 2023-2024 the original author or 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
*
* https://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 org.springframework.grpc.client;
import java.time.Duration;
import java.util.function.Consumer;
import org.springframework.util.unit.DataSize;
import io.grpc.ManagedChannel;
/**
* Represents the configuration for a {@link ManagedChannel gRPC channel}.
*
* @author Chris Bono
*/
public class NamedChannel {
/**
* The target address uri to connect to.
*/
private String address = "static://localhost:9090";
public String getAddress() {
return this.address;
}
public void setAddress(final String address) {
this.address = address;
}
// --------------------------------------------------
// defaultLoadBalancingPolicy
// --------------------------------------------------
/**
* The default load balancing policy the channel should use.
*/
private String defaultLoadBalancingPolicy = "round_robin";
public String getDefaultLoadBalancingPolicy() {
return this.defaultLoadBalancingPolicy;
}
public void setDefaultLoadBalancingPolicy(final String defaultLoadBalancingPolicy) {
this.defaultLoadBalancingPolicy = defaultLoadBalancingPolicy;
}
// --------------------------------------------------
private final Health health = new Health();
public Health getHealth() {
return this.health;
}
/**
* The negotiation type for the channel.
*/
private NegotiationType negotiationType = NegotiationType.PLAINTEXT;
public NegotiationType getNegotiationType() {
return this.negotiationType;
}
public void setNegotiationType(NegotiationType negotiationType) {
this.negotiationType = negotiationType;
}
// --------------------------------------------------
// KeepAlive
// --------------------------------------------------
/**
* Whether keep alive is enabled on the channel.
*/
private boolean enableKeepAlive = false;
public boolean isEnableKeepAlive() {
return this.enableKeepAlive;
}
public void setEnableKeepAlive(boolean enableKeepAlive) {
this.enableKeepAlive = enableKeepAlive;
}
// --------------------------------------------------
/**
* The duration without ongoing RPCs before going to idle mode.
*/
private Duration idleTimeout = Duration.ofSeconds(20);
public Duration getIdleTimeout() {
return this.idleTimeout;
}
public void setIdleTimeout(Duration idleTimeout) {
this.idleTimeout = idleTimeout;
}
// --------------------------------------------------
/**
* The delay before sending a keepAlive. Note that shorter intervals increase the
* network burden for the server and this value can not be lower than
* 'permitKeepAliveTime' on the server.
*/
private Duration keepAliveTime = Duration.ofMinutes(5);
public Duration getKeepAliveTime() {
return this.keepAliveTime;
}
public void setKeepAliveTime(Duration keepAliveTime) {
this.keepAliveTime = keepAliveTime;
}
// --------------------------------------------------
/**
* The default timeout for a keepAlives ping request.
*/
private Duration keepAliveTimeout = Duration.ofSeconds(20);
public Duration getKeepAliveTimeout() {
return this.keepAliveTimeout;
}
public void setKeepAliveTimeout(Duration keepAliveTimeout) {
this.keepAliveTimeout = keepAliveTimeout;
}
// --------------------------------------------------
/**
* Whether a keepAlive will be performed when there are no outstanding RPC on a
* connection.
*/
private boolean keepAliveWithoutCalls = false;
public boolean isKeepAliveWithoutCalls() {
return this.keepAliveWithoutCalls;
}
public void setKeepAliveWithoutCalls(boolean keepAliveWithoutCalls) {
this.keepAliveWithoutCalls = keepAliveWithoutCalls;
}
// --------------------------------------------------
// Message Transfer
// --------------------------------------------------
/**
* Maximum message size allowed to be received by the channel (default 4MiB). Set to
* '-1' to use the highest possible limit (not recommended).
*/
private DataSize maxInboundMessageSize = DataSize.ofBytes(4194304);
/**
* Maximum metadata size allowed to be received by the channel (default 8KiB). Set to
* '-1' to use the highest possible limit (not recommended).
*/
private DataSize maxInboundMetadataSize = DataSize.ofBytes(8192);
public DataSize getMaxInboundMessageSize() {
return this.maxInboundMessageSize;
}
public void setMaxInboundMessageSize(final DataSize maxInboundMessageSize) {
this.setMaxInboundSize(maxInboundMessageSize, (s) -> this.maxInboundMessageSize = s, "maxInboundMesssageSize");
}
public DataSize getMaxInboundMetadataSize() {
return this.maxInboundMetadataSize;
}
public void setMaxInboundMetadataSize(DataSize maxInboundMetadataSize) {
this.setMaxInboundSize(maxInboundMetadataSize, (s) -> this.maxInboundMetadataSize = s,
"maxInboundMetadataSize");
}
private void setMaxInboundSize(DataSize maxSize, Consumer<DataSize> setter, String propertyName) {
if (maxSize != null && maxSize.toBytes() >= 0) {
setter.accept(maxSize);
}
else if (maxSize != null && maxSize.toBytes() == -1) {
setter.accept(DataSize.ofBytes(Integer.MAX_VALUE));
}
else {
throw new IllegalArgumentException("Unsupported %s: %s".formatted(propertyName, maxSize));
}
}
// --------------------------------------------------
/**
* The custom User-Agent for the channel.
*/
private String userAgent = null;
public String getUserAgent() {
return this.userAgent;
}
public void setUserAgent(final String userAgent) {
this.userAgent = userAgent;
}
/**
* Provide a copy of the channel instance.
* @return a copy of the channel instance.
*/
public NamedChannel copy() {
NamedChannel copy = new NamedChannel();
copy.address = this.address;
copy.defaultLoadBalancingPolicy = this.defaultLoadBalancingPolicy;
copy.negotiationType = this.negotiationType;
copy.enableKeepAlive = this.enableKeepAlive;
copy.idleTimeout = this.idleTimeout;
copy.keepAliveTime = this.keepAliveTime;
copy.keepAliveTimeout = this.keepAliveTimeout;
copy.keepAliveWithoutCalls = this.keepAliveWithoutCalls;
copy.maxInboundMessageSize = this.maxInboundMessageSize;
copy.maxInboundMetadataSize = this.maxInboundMetadataSize;
copy.userAgent = this.userAgent;
copy.health.copyValuesFrom(this.getHealth());
copy.ssl.copyValuesFrom(this.getSsl());
return copy;
}
// --------------------------------------------------
/**
* Flag to say that strict SSL checks are not enabled (so the remote certificate could
* be anonymous).
*/
private boolean secure = true;
public boolean isSecure() {
return this.secure;
}
public void setSecure(boolean secure) {
this.secure = secure;
}
// --------------------------------------------------
private final Ssl ssl = new Ssl();
public Ssl getSsl() {
return this.ssl;
}
public static class Ssl {
/**
* Whether to enable SSL support. Enabled automatically if "bundle" is provided
* unless specified otherwise.
*/
private Boolean enabled;
/**
* SSL bundle name.
*/
private String bundle;
public boolean isEnabled() {
return (this.enabled != null) ? this.enabled : this.bundle != null;
}
public void setEnabled(boolean enabled) {
this.enabled = enabled;
}
public String getBundle() {
return this.bundle;
}
public void setBundle(String bundle) {
this.bundle = bundle;
}
/**
* Copies the values from another instance.
* @param other instance to copy values from
*/
public void copyValuesFrom(Ssl other) {
this.enabled = other.enabled;
this.bundle = other.bundle;
}
}
public static class Health {
/**
* Whether to enable client-side health check for the channel.
*/
private boolean enabled = false;
/**
* Name of the service to check health on.
*/
private String serviceName;
public boolean isEnabled() {
return this.enabled;
}
public void setEnabled(boolean enabled) {
this.enabled = enabled;
}
public String getServiceName() {
return this.serviceName;
}
public void setServiceName(String serviceName) {
this.serviceName = serviceName;
}
/**
* Copies the values from another instance.
* @param other instance to copy values from
*/
public void copyValuesFrom(Health other) {
this.enabled = other.enabled;
this.serviceName = other.serviceName;
}
}
}

View File

@@ -1,88 +0,0 @@
/*
* Copyright 2023-2024 the original author or 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
*
* https://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 org.springframework.grpc.client;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import org.springframework.context.EnvironmentAware;
import org.springframework.core.env.Environment;
import org.springframework.core.env.StandardEnvironment;
/**
* Provides access to configured {@link NamedChannel channels}.
*
* @author Chris Bono
*/
public class NamedChannelRegistry implements EnvironmentAware, VirtualTargets {
private final NamedChannel defaultChannel;
private final Map<String, NamedChannel> channels;
private Environment environment;
public NamedChannelRegistry(NamedChannel defaultChannel, Map<String, NamedChannel> configuredChannels) {
this.defaultChannel = defaultChannel;
this.channels = new ConcurrentHashMap<>(configuredChannels);
this.environment = new StandardEnvironment();
}
@Override
public void setEnvironment(Environment environment) {
this.environment = environment;
}
/**
* Gets the default channel configuration.
* @return the default channel
*/
public NamedChannel getDefaultChannel() {
return this.defaultChannel;
}
/**
* Gets the configured channel with the given name. If no channel is configured for
* the specified name then one is created using the default channel as a template.
* @param name the name of the channel
* @return the configured channel if found, or a newly created channel using the
* default channel as a template
*/
public NamedChannel getChannel(String name) {
if ("default".equals(name)) {
return this.defaultChannel;
}
return this.channels.computeIfAbsent(name, authority -> {
NamedChannel channel = this.defaultChannel.copy();
if (!authority.contains(":/") && !authority.startsWith("unix:")) {
authority = "static://" + authority;
}
channel.setAddress(authority);
return channel;
});
}
@Override
public String getTarget(String authority) {
NamedChannel channel = this.getChannel(authority);
String address = channel.getAddress();
if (address.startsWith("static:") || address.startsWith("tcp:")) {
address = address.substring(address.indexOf(":") + 1).replaceFirst("/*", "");
}
return this.environment.resolvePlaceholders(address);
}
}

View File

@@ -1,138 +0,0 @@
/*
* Copyright 2023-2024 the original author or 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
*
* https://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 org.springframework.grpc.client;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.entry;
import java.time.Duration;
import java.util.Map;
import org.assertj.core.api.InstanceOfAssertFactories;
import org.junit.jupiter.api.Nested;
import org.junit.jupiter.api.Test;
import org.springframework.mock.env.MockEnvironment;
import org.springframework.util.unit.DataSize;
/**
* Tests for {@link NamedChannelRegistry}.
*/
class NamedChannelRegistryTests {
@Nested
class GetChannelAPI {
@Test
void withDefaultNameReturnsDefaultChannel() {
var defaultChannel = new NamedChannel();
var registry = new NamedChannelRegistry(defaultChannel, Map.of());
assertThat(registry.getChannel("default")).isSameAs(defaultChannel);
assertThat(registry.getDefaultChannel()).isSameAs(defaultChannel);
assertThat(registry).extracting("channels", InstanceOfAssertFactories.MAP).isEmpty();
}
@Test
void withKnownNameReturnsKnownChannel() {
var defaultChannel = new NamedChannel();
var channel1 = new NamedChannel();
var registry = new NamedChannelRegistry(defaultChannel, Map.of("c1", channel1));
assertThat(registry.getChannel("c1")).isSameAs(channel1);
assertThat(registry).extracting("channels", InstanceOfAssertFactories.MAP)
.containsExactly(entry("c1", channel1));
}
@Test
void withUnknownNameReturnsNewChannelWithCopiedDefaults() {
var defaultChannel = new NamedChannel();
defaultChannel.setAddress("static://my-server:9999");
defaultChannel.setDefaultLoadBalancingPolicy("custom");
defaultChannel.getHealth().setEnabled(true);
defaultChannel.getHealth().setServiceName("custom-service");
defaultChannel.setEnableKeepAlive(true);
defaultChannel.setIdleTimeout(Duration.ofMinutes(1));
defaultChannel.setKeepAliveTime(Duration.ofMinutes(4));
defaultChannel.setKeepAliveTimeout(Duration.ofMinutes(6));
defaultChannel.setKeepAliveWithoutCalls(true);
defaultChannel.setMaxInboundMessageSize(DataSize.ofMegabytes(100));
defaultChannel.setMaxInboundMetadataSize(DataSize.ofMegabytes(200));
defaultChannel.setUserAgent("me");
defaultChannel.getSsl().setEnabled(true);
defaultChannel.getSsl().setBundle("custom-bundle");
var registry = new NamedChannelRegistry(defaultChannel, Map.of());
var newChannel = registry.getChannel("new-channel");
assertThat(newChannel).usingRecursiveComparison().ignoringFields("address").isEqualTo(defaultChannel);
assertThat(registry).extracting("channels", InstanceOfAssertFactories.MAP)
.containsExactly(entry("new-channel", newChannel));
}
}
@Nested
class GetTargetAPI {
@Test
void channelWithStaticAddressReturnsStrippedAddress() {
var defaultChannel = new NamedChannel();
var channel1 = new NamedChannel();
channel1.setAddress("static://my-server:8888");
var registry = new NamedChannelRegistry(defaultChannel, Map.of("c1", channel1));
assertThat(registry.getTarget("c1")).isEqualTo("my-server:8888");
assertThat(registry).extracting("channels", InstanceOfAssertFactories.MAP)
.containsExactly(entry("c1", channel1));
}
@Test
void channelWithTcpAddressReturnsStrippedAddress() {
var defaultChannel = new NamedChannel();
var channel1 = new NamedChannel();
channel1.setAddress("tcp://my-server:8888");
var registry = new NamedChannelRegistry(defaultChannel, Map.of("c1", channel1));
assertThat(registry.getTarget("c1")).isEqualTo("my-server:8888");
assertThat(registry).extracting("channels", InstanceOfAssertFactories.MAP)
.containsExactly(entry("c1", channel1));
}
@Test
void channelWithAddressPropertyPlaceholdersPopulatesFromEnvironment() {
var defaultChannel = new NamedChannel();
var channel1 = new NamedChannel();
channel1.setAddress("my-server-${channelName}:8888");
var registry = new NamedChannelRegistry(defaultChannel, Map.of("c1", channel1));
var env = new MockEnvironment();
env.setProperty("channelName", "foo");
registry.setEnvironment(env);
assertThat(registry.getTarget("c1")).isEqualTo("my-server-foo:8888");
assertThat(registry).extracting("channels", InstanceOfAssertFactories.MAP)
.containsExactly(entry("c1", channel1));
}
}
@Nested
class CopyDefaultsAPI {
@Test
void copyFromDefaultChannel() {
var registry = new NamedChannelRegistry(new NamedChannel(), Map.of());
var defaultChannel = registry.getDefaultChannel();
var newChannel = defaultChannel.copy();
assertThat(newChannel).usingRecursiveComparison().isEqualTo(defaultChannel);
}
}
}

View File

@@ -1,7 +1,7 @@
|===
|Name | Default | Description
|spring.grpc.client.channels | |
|spring.grpc.client.channels | | Map of channels configured by name.
|spring.grpc.client.default-channel.address | `+++static://localhost:9090+++` | The target address uri to connect to.
|spring.grpc.client.default-channel.default-load-balancing-policy | `+++round_robin+++` | The default load balancing policy the channel should use.
|spring.grpc.client.default-channel.enable-keep-alive | `+++false+++` | Whether keep alive is enabled on the channel.
@@ -46,4 +46,4 @@
|spring.grpc.server.ssl.enabled | | Whether to enable SSL support. Enabled automatically if "bundle" is provided unless specified otherwise.
|spring.grpc.server.ssl.secure | `+++true+++` | Flag to indicate that client authentication is secure (i.e. certificates are checked). Do not set this to false in production.
|===
|===

View File

@@ -22,9 +22,8 @@ import java.util.function.BiConsumer;
import java.util.function.Consumer;
import org.springframework.boot.context.properties.PropertyMapper;
import org.springframework.grpc.autoconfigure.client.GrpcClientProperties.ChannelConfig;
import org.springframework.grpc.client.GrpcChannelBuilderCustomizer;
import org.springframework.grpc.client.NamedChannel;
import org.springframework.grpc.client.NamedChannelRegistry;
import org.springframework.util.unit.DataSize;
import io.grpc.ManagedChannelBuilder;
@@ -40,18 +39,15 @@ import io.grpc.ManagedChannelBuilder;
class ClientPropertiesChannelBuilderCustomizer<T extends ManagedChannelBuilder<T>>
implements GrpcChannelBuilderCustomizer<T> {
private final NamedChannelRegistry channelRegistry;
private final GrpcClientProperties properties;
ClientPropertiesChannelBuilderCustomizer(NamedChannelRegistry channelRegistry) {
this.channelRegistry = channelRegistry;
ClientPropertiesChannelBuilderCustomizer(GrpcClientProperties properties) {
this.properties = properties;
}
@Override
public void customize(String authority, T builder) {
NamedChannel channel = this.channelRegistry.getChannel(authority);
if (channel == null) {
return;
}
ChannelConfig channel = this.properties.getChannel(authority);
PropertyMapper mapper = PropertyMapper.get().alwaysApplyingWhenNonNull();
mapper.from(channel.getUserAgent()).to(builder::userAgent);
if (!authority.startsWith("unix:")) {

View File

@@ -27,7 +27,6 @@ import org.springframework.grpc.client.ChannelCredentialsProvider;
import org.springframework.grpc.client.ClientInterceptorsConfigurer;
import org.springframework.grpc.client.GrpcChannelBuilderCustomizer;
import org.springframework.grpc.client.GrpcChannelFactory;
import org.springframework.grpc.client.NamedChannelRegistry;
import org.springframework.grpc.client.NettyGrpcChannelFactory;
import org.springframework.grpc.client.ShadedNettyGrpcChannelFactory;
@@ -47,14 +46,14 @@ class GrpcChannelFactoryConfigurations {
static class ShadedNettyChannelFactoryConfiguration {
@Bean
ShadedNettyGrpcChannelFactory shadedNettyGrpcChannelFactory(NamedChannelRegistry namedChannelRegistry,
ShadedNettyGrpcChannelFactory shadedNettyGrpcChannelFactory(GrpcClientProperties properties,
ChannelBuilderCustomizers channelBuilderCustomizers,
ClientInterceptorsConfigurer interceptorsConfigurer, ChannelCredentialsProvider credentials) {
List<GrpcChannelBuilderCustomizer<io.grpc.netty.shaded.io.grpc.netty.NettyChannelBuilder>> builderCustomizers = List
.of(channelBuilderCustomizers::customize);
var factory = new ShadedNettyGrpcChannelFactory(builderCustomizers, interceptorsConfigurer);
factory.setCredentialsProvider(credentials);
factory.setVirtualTargets(namedChannelRegistry);
factory.setVirtualTargets(properties);
return factory;
}
@@ -67,14 +66,14 @@ class GrpcChannelFactoryConfigurations {
static class NettyChannelFactoryConfiguration {
@Bean
NettyGrpcChannelFactory nettyGrpcChannelFactory(NamedChannelRegistry namedChannelRegistry,
NettyGrpcChannelFactory nettyGrpcChannelFactory(GrpcClientProperties properties,
ChannelBuilderCustomizers channelBuilderCustomizers,
ClientInterceptorsConfigurer interceptorsConfigurer, ChannelCredentialsProvider credentials) {
List<GrpcChannelBuilderCustomizer<NettyChannelBuilder>> builderCustomizers = List
.of(channelBuilderCustomizers::customize);
var factory = new NettyGrpcChannelFactory(builderCustomizers, interceptorsConfigurer);
factory.setCredentialsProvider(credentials);
factory.setVirtualTargets(namedChannelRegistry);
factory.setVirtualTargets(properties);
return factory;
}

View File

@@ -26,7 +26,6 @@ import org.springframework.context.annotation.Import;
import org.springframework.grpc.autoconfigure.common.codec.GrpcCodecConfiguration;
import org.springframework.grpc.client.ChannelCredentialsProvider;
import org.springframework.grpc.client.GrpcChannelBuilderCustomizer;
import org.springframework.grpc.client.NamedChannelRegistry;
import io.grpc.CompressorRegistry;
import io.grpc.DecompressorRegistry;
@@ -39,23 +38,16 @@ import io.grpc.ManagedChannelBuilder;
GrpcChannelFactoryConfigurations.NettyChannelFactoryConfiguration.class })
public class GrpcClientAutoConfiguration {
@Bean
@ConditionalOnMissingBean
NamedChannelRegistry namedChannelRegistry(GrpcClientProperties properties) {
return new NamedChannelRegistry(properties.getDefaultChannel(), properties.getChannels());
}
@Bean
@ConditionalOnMissingBean(ChannelCredentialsProvider.class)
NamedChannelCredentialsProvider channelCredentialsProvider(SslBundles bundles,
NamedChannelRegistry channelRegistry) {
return new NamedChannelCredentialsProvider(bundles, channelRegistry);
NamedChannelCredentialsProvider channelCredentialsProvider(SslBundles bundles, GrpcClientProperties properties) {
return new NamedChannelCredentialsProvider(bundles, properties);
}
@Bean
<T extends ManagedChannelBuilder<T>> GrpcChannelBuilderCustomizer<T> clientPropertiesChannelCustomizer(
NamedChannelRegistry channelRegistry) {
return new ClientPropertiesChannelBuilderCustomizer<>(channelRegistry);
GrpcClientProperties properties) {
return new ClientPropertiesChannelBuilderCustomizer<>(properties);
}
@ConditionalOnBean(CompressorRegistry.class)

View File

@@ -15,35 +15,415 @@
*/
package org.springframework.grpc.autoconfigure.client;
import java.time.Duration;
import java.time.temporal.ChronoUnit;
import java.util.HashMap;
import java.util.Map;
import java.util.function.Consumer;
import org.springframework.boot.context.properties.ConfigurationProperties;
import org.springframework.boot.context.properties.NestedConfigurationProperty;
import org.springframework.grpc.client.NamedChannel;
import org.springframework.boot.convert.DurationUnit;
import org.springframework.context.EnvironmentAware;
import org.springframework.core.env.Environment;
import org.springframework.core.env.StandardEnvironment;
import org.springframework.grpc.client.NegotiationType;
import org.springframework.grpc.client.VirtualTargets;
import org.springframework.util.unit.DataSize;
import io.grpc.ManagedChannel;
@ConfigurationProperties(prefix = "spring.grpc.client")
public class GrpcClientProperties {
public class GrpcClientProperties implements EnvironmentAware, VirtualTargets {
@NestedConfigurationProperty
private final NamedChannel defaultChannel = new NamedChannel();
/**
* The default channel configuration to use for new channels.
*/
private final ChannelConfig defaultChannel = new ChannelConfig();
private final Map<String, NamedChannel> channels = new HashMap<>();
/**
* Map of channels configured by name.
*/
private final Map<String, ChannelConfig> channels = new HashMap<>();
private Environment environment;
GrpcClientProperties() {
this.defaultChannel.setAddress("static://localhost:9090");
this.environment = new StandardEnvironment();
}
public Map<String, NamedChannel> getChannels() {
return this.channels;
}
/**
* Gets the default {@link NamedChannel} configured for the GRPC client.
* @return the default {@link NamedChannel}
*/
public NamedChannel getDefaultChannel() {
public ChannelConfig getDefaultChannel() {
return this.defaultChannel;
}
public Map<String, ChannelConfig> getChannels() {
return this.channels;
}
@Override
public void setEnvironment(Environment environment) {
this.environment = environment;
}
/**
* Gets the configured channel with the given name. If no channel is configured for
* the specified name then one is created using the default channel as a template.
* @param name the name of the channel
* @return the configured channel if found, or a newly created channel using the
* default channel as a template
*/
public ChannelConfig getChannel(String name) {
if ("default".equals(name)) {
return this.defaultChannel;
}
ChannelConfig channel = this.channels.get(name);
if (channel != null) {
return channel;
}
channel = this.defaultChannel.copy();
String address = name;
if (!name.contains(":/") && !name.startsWith("unix:")) {
address = "static://" + name;
}
channel.setAddress(address);
return channel;
}
@Override
public String getTarget(String authority) {
ChannelConfig channel = this.getChannel(authority);
String address = channel.getAddress();
if (address.startsWith("static:") || address.startsWith("tcp:")) {
address = address.substring(address.indexOf(":") + 1).replaceFirst("/*", "");
}
return this.environment.resolvePlaceholders(address);
}
/**
* Represents the configuration for a {@link ManagedChannel gRPC channel}.
*/
public static class ChannelConfig {
/**
* The target address uri to connect to.
*/
private String address = "static://localhost:9090";
public String getAddress() {
return this.address;
}
public void setAddress(final String address) {
this.address = address;
}
// --------------------------------------------------
// defaultLoadBalancingPolicy
// --------------------------------------------------
/**
* The default load balancing policy the channel should use.
*/
private String defaultLoadBalancingPolicy = "round_robin";
public String getDefaultLoadBalancingPolicy() {
return this.defaultLoadBalancingPolicy;
}
public void setDefaultLoadBalancingPolicy(final String defaultLoadBalancingPolicy) {
this.defaultLoadBalancingPolicy = defaultLoadBalancingPolicy;
}
// --------------------------------------------------
private final Health health = new Health();
public Health getHealth() {
return this.health;
}
/**
* The negotiation type for the channel.
*/
private NegotiationType negotiationType = NegotiationType.PLAINTEXT;
public NegotiationType getNegotiationType() {
return this.negotiationType;
}
public void setNegotiationType(NegotiationType negotiationType) {
this.negotiationType = negotiationType;
}
// --------------------------------------------------
// KeepAlive
// --------------------------------------------------
/**
* Whether keep alive is enabled on the channel.
*/
private boolean enableKeepAlive = false;
public boolean isEnableKeepAlive() {
return this.enableKeepAlive;
}
public void setEnableKeepAlive(boolean enableKeepAlive) {
this.enableKeepAlive = enableKeepAlive;
}
// --------------------------------------------------
/**
* The duration without ongoing RPCs before going to idle mode.
*/
@DurationUnit(ChronoUnit.SECONDS)
private Duration idleTimeout = Duration.ofSeconds(20);
public Duration getIdleTimeout() {
return this.idleTimeout;
}
public void setIdleTimeout(Duration idleTimeout) {
this.idleTimeout = idleTimeout;
}
// --------------------------------------------------
/**
* The delay before sending a keepAlive. Note that shorter intervals increase the
* network burden for the server and this value can not be lower than
* 'permitKeepAliveTime' on the server.
*/
@DurationUnit(ChronoUnit.SECONDS)
private Duration keepAliveTime = Duration.ofMinutes(5);
public Duration getKeepAliveTime() {
return this.keepAliveTime;
}
public void setKeepAliveTime(Duration keepAliveTime) {
this.keepAliveTime = keepAliveTime;
}
// --------------------------------------------------
/**
* The default timeout for a keepAlives ping request.
*/
@DurationUnit(ChronoUnit.SECONDS)
private Duration keepAliveTimeout = Duration.ofSeconds(20);
public Duration getKeepAliveTimeout() {
return this.keepAliveTimeout;
}
public void setKeepAliveTimeout(Duration keepAliveTimeout) {
this.keepAliveTimeout = keepAliveTimeout;
}
// --------------------------------------------------
/**
* Whether a keepAlive will be performed when there are no outstanding RPC on a
* connection.
*/
private boolean keepAliveWithoutCalls = false;
public boolean isKeepAliveWithoutCalls() {
return this.keepAliveWithoutCalls;
}
public void setKeepAliveWithoutCalls(boolean keepAliveWithoutCalls) {
this.keepAliveWithoutCalls = keepAliveWithoutCalls;
}
// --------------------------------------------------
// Message Transfer
// --------------------------------------------------
/**
* Maximum message size allowed to be received by the channel (default 4MiB). Set
* to '-1' to use the highest possible limit (not recommended).
*/
private DataSize maxInboundMessageSize = DataSize.ofBytes(4194304);
/**
* Maximum metadata size allowed to be received by the channel (default 8KiB). Set
* to '-1' to use the highest possible limit (not recommended).
*/
private DataSize maxInboundMetadataSize = DataSize.ofBytes(8192);
public DataSize getMaxInboundMessageSize() {
return this.maxInboundMessageSize;
}
public void setMaxInboundMessageSize(final DataSize maxInboundMessageSize) {
this.setMaxInboundSize(maxInboundMessageSize, (s) -> this.maxInboundMessageSize = s,
"maxInboundMesssageSize");
}
public DataSize getMaxInboundMetadataSize() {
return this.maxInboundMetadataSize;
}
public void setMaxInboundMetadataSize(DataSize maxInboundMetadataSize) {
this.setMaxInboundSize(maxInboundMetadataSize, (s) -> this.maxInboundMetadataSize = s,
"maxInboundMetadataSize");
}
private void setMaxInboundSize(DataSize maxSize, Consumer<DataSize> setter, String propertyName) {
if (maxSize != null && maxSize.toBytes() >= 0) {
setter.accept(maxSize);
}
else if (maxSize != null && maxSize.toBytes() == -1) {
setter.accept(DataSize.ofBytes(Integer.MAX_VALUE));
}
else {
throw new IllegalArgumentException("Unsupported %s: %s".formatted(propertyName, maxSize));
}
}
// --------------------------------------------------
/**
* The custom User-Agent for the channel.
*/
private String userAgent = null;
public String getUserAgent() {
return this.userAgent;
}
public void setUserAgent(final String userAgent) {
this.userAgent = userAgent;
}
/**
* Provide a copy of the channel instance.
* @return a copy of the channel instance.
*/
public ChannelConfig copy() {
ChannelConfig copy = new ChannelConfig();
copy.address = this.address;
copy.defaultLoadBalancingPolicy = this.defaultLoadBalancingPolicy;
copy.negotiationType = this.negotiationType;
copy.enableKeepAlive = this.enableKeepAlive;
copy.idleTimeout = this.idleTimeout;
copy.keepAliveTime = this.keepAliveTime;
copy.keepAliveTimeout = this.keepAliveTimeout;
copy.keepAliveWithoutCalls = this.keepAliveWithoutCalls;
copy.maxInboundMessageSize = this.maxInboundMessageSize;
copy.maxInboundMetadataSize = this.maxInboundMetadataSize;
copy.userAgent = this.userAgent;
copy.health.copyValuesFrom(this.getHealth());
copy.ssl.copyValuesFrom(this.getSsl());
return copy;
}
// --------------------------------------------------
/**
* Flag to say that strict SSL checks are not enabled (so the remote certificate
* could be anonymous).
*/
private boolean secure = true;
public boolean isSecure() {
return this.secure;
}
public void setSecure(boolean secure) {
this.secure = secure;
}
// --------------------------------------------------
private final Ssl ssl = new Ssl();
public Ssl getSsl() {
return this.ssl;
}
public static class Ssl {
/**
* Whether to enable SSL support. Enabled automatically if "bundle" is
* provided unless specified otherwise.
*/
private Boolean enabled;
/**
* SSL bundle name.
*/
private String bundle;
public boolean isEnabled() {
return (this.enabled != null) ? this.enabled : this.bundle != null;
}
public void setEnabled(boolean enabled) {
this.enabled = enabled;
}
public String getBundle() {
return this.bundle;
}
public void setBundle(String bundle) {
this.bundle = bundle;
}
/**
* Copies the values from another instance.
* @param other instance to copy values from
*/
public void copyValuesFrom(Ssl other) {
this.enabled = other.enabled;
this.bundle = other.bundle;
}
}
public static class Health {
/**
* Whether to enable client-side health check for the channel.
*/
private boolean enabled = false;
/**
* Name of the service to check health on.
*/
private String serviceName;
public boolean isEnabled() {
return this.enabled;
}
public void setEnabled(boolean enabled) {
this.enabled = enabled;
}
public String getServiceName() {
return this.serviceName;
}
public void setServiceName(String serviceName) {
this.serviceName = serviceName;
}
/**
* Copies the values from another instance.
* @param other instance to copy values from
*/
public void copyValuesFrom(Health other) {
this.enabled = other.enabled;
this.serviceName = other.serviceName;
}
}
}
}

View File

@@ -19,9 +19,8 @@ import javax.net.ssl.TrustManagerFactory;
import org.springframework.boot.ssl.SslBundle;
import org.springframework.boot.ssl.SslBundles;
import org.springframework.grpc.autoconfigure.client.GrpcClientProperties.ChannelConfig;
import org.springframework.grpc.client.ChannelCredentialsProvider;
import org.springframework.grpc.client.NamedChannel;
import org.springframework.grpc.client.NamedChannelRegistry;
import org.springframework.grpc.client.NegotiationType;
import org.springframework.grpc.internal.InsecureTrustManagerFactory;
@@ -29,20 +28,25 @@ import io.grpc.ChannelCredentials;
import io.grpc.InsecureChannelCredentials;
import io.grpc.TlsChannelCredentials;
/**
* Provides channel credentials using channel configuration and {@link SslBundles}.
*
* @author David Syer
*/
public class NamedChannelCredentialsProvider implements ChannelCredentialsProvider {
private final SslBundles bundles;
private final NamedChannelRegistry channelRegistry;
private final GrpcClientProperties properties;
public NamedChannelCredentialsProvider(SslBundles bundles, NamedChannelRegistry channelRegistry) {
public NamedChannelCredentialsProvider(SslBundles bundles, GrpcClientProperties properties) {
this.bundles = bundles;
this.channelRegistry = channelRegistry;
this.properties = properties;
}
@Override
public ChannelCredentials getChannelCredentials(String path) {
NamedChannel channel = this.channelRegistry.getChannel(path);
ChannelConfig channel = this.properties.getChannel(path);
if (!channel.getSsl().isEnabled() && channel.getNegotiationType() == NegotiationType.PLAINTEXT) {
return InsecureChannelCredentials.create();
}

View File

@@ -43,8 +43,6 @@ import org.springframework.core.annotation.Order;
import org.springframework.grpc.client.ChannelCredentialsProvider;
import org.springframework.grpc.client.GrpcChannelBuilderCustomizer;
import org.springframework.grpc.client.GrpcChannelFactory;
import org.springframework.grpc.client.NamedChannel;
import org.springframework.grpc.client.NamedChannelRegistry;
import org.springframework.grpc.client.NettyGrpcChannelFactory;
import org.springframework.grpc.client.ShadedNettyGrpcChannelFactory;
@@ -67,26 +65,6 @@ class GrpcClientAutoConfigurationTests {
.withConfiguration(AutoConfigurations.of(GrpcClientAutoConfiguration.class, SslAutoConfiguration.class));
}
@Test
void whenHasUserDefinedChannelRegistryDoesNotAutoConfigureBean() {
NamedChannelRegistry customChannelRegistry = mock(NamedChannelRegistry.class);
this.contextRunner()
.withBean("customChannelRegistry", NamedChannelRegistry.class, () -> customChannelRegistry)
.run((context) -> assertThat(context).getBean(NamedChannelRegistry.class).isSameAs(customChannelRegistry));
}
@Test
void channelRegistryAutoConfiguredAsExpected() {
this.contextRunner()
.run((context) -> assertThat(context).getBean(NamedChannelRegistry.class).satisfies((channelRegistry) -> {
var properties = context.getBean(GrpcClientProperties.class);
assertThat(channelRegistry.getDefaultChannel()).isEqualTo(properties.getDefaultChannel());
assertThat(channelRegistry)
.extracting("channels", InstanceOfAssertFactories.map(String.class, NamedChannel.class))
.containsExactlyInAnyOrderEntriesOf(properties.getChannels());
}));
}
@Test
void whenHasUserDefinedCredentialsProviderDoesNotAutoConfigureBean() {
ChannelCredentialsProvider customCredentialsProvider = mock(ChannelCredentialsProvider.class);
@@ -100,11 +78,20 @@ class GrpcClientAutoConfigurationTests {
void credentialsProviderAutoConfiguredAsExpected() {
this.contextRunner()
.run((context) -> assertThat(context).getBean(NamedChannelCredentialsProvider.class)
.hasFieldOrPropertyWithValue("channelRegistry", context.getBean(NamedChannelRegistry.class))
.hasFieldOrPropertyWithValue("properties", context.getBean(GrpcClientProperties.class))
.extracting("bundles")
.isInstanceOf(SslBundles.class));
}
@Test
void clientPropertiesAutoConfiguredResolvesPlaceholders() {
this.contextRunner()
.withPropertyValues("spring.grpc.client.channels.c1.address=my-server-${channelName}:8888",
"channelName=foo")
.run((context) -> assertThat(context).getBean(GrpcClientProperties.class)
.satisfies((properties) -> assertThat(properties.getTarget("c1")).isEqualTo("my-server-foo:8888")));
}
@Test
void clientPropertiesChannelCustomizerAutoConfiguredWithHealthAsExpected() {
this.contextRunner()
@@ -249,7 +236,7 @@ class GrpcClientAutoConfigurationTests {
.isInstanceOf(expectedChannelFactoryType)
.hasFieldOrPropertyWithValue("credentials", context.getBean(NamedChannelCredentialsProvider.class))
.extracting("targets")
.isInstanceOf(NamedChannelRegistry.class));
.isInstanceOf(GrpcClientProperties.class));
}
@Test

View File

@@ -17,19 +17,23 @@
package org.springframework.grpc.autoconfigure.client;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.entry;
import java.time.Duration;
import java.util.HashMap;
import java.util.Map;
import java.util.function.Function;
import org.assertj.core.api.InstanceOfAssertFactories;
import org.junit.jupiter.api.Nested;
import org.junit.jupiter.api.Test;
import org.springframework.boot.context.properties.bind.Binder;
import org.springframework.boot.context.properties.source.MapConfigurationPropertySource;
import org.springframework.grpc.client.NamedChannel;
import org.springframework.grpc.autoconfigure.client.GrpcClientProperties.ChannelConfig;
import org.springframework.grpc.client.NegotiationType;
import org.springframework.mock.env.MockEnvironment;
import org.springframework.test.util.ReflectionTestUtils;
import org.springframework.util.unit.DataSize;
/**
@@ -45,6 +49,13 @@ class GrpcClientPropertiesTests {
.get();
}
private GrpcClientProperties newProperties(ChannelConfig defaultChannel, Map<String, ChannelConfig> channels) {
var properties = new GrpcClientProperties();
ReflectionTestUtils.setField(properties, "defaultChannel", defaultChannel);
ReflectionTestUtils.setField(properties, "channels", channels);
return properties;
}
@Nested
class BindPropertiesAPI {
@@ -55,11 +66,11 @@ class GrpcClientPropertiesTests {
@Test
void specificChannelWithDefaultValues() {
this.withDefaultValues("channels.c1", (p) -> p.getChannels().get("c1"));
this.withDefaultValues("channels.c1", (p) -> p.getChannel("c1"));
}
private void withDefaultValues(String channelName,
Function<GrpcClientProperties, NamedChannel> channelFromProperties) {
Function<GrpcClientProperties, ChannelConfig> channelFromProperties) {
Map<String, String> map = new HashMap<>();
// we have to at least bind one property or bind() fails
map.put("spring.grpc.client.%s.enable-keep-alive".formatted(channelName), "false");
@@ -91,11 +102,11 @@ class GrpcClientPropertiesTests {
@Test
void specificChannelWithSpecifiedValues() {
this.withSpecifiedValues("channels.c1", (p) -> p.getChannels().get("c1"));
this.withSpecifiedValues("channels.c1", (p) -> p.getChannel("c1"));
}
private void withSpecifiedValues(String channelName,
Function<GrpcClientProperties, NamedChannel> channelFromProperties) {
Function<GrpcClientProperties, ChannelConfig> channelFromProperties) {
Map<String, String> map = new HashMap<>();
var propPrefix = "spring.grpc.client.%s.".formatted(channelName);
map.put("%s.address".formatted(propPrefix), "static://my-server:8888");
@@ -138,9 +149,9 @@ class GrpcClientPropertiesTests {
@Test
void withoutKeepAliveUnitsSpecified() {
Map<String, String> map = new HashMap<>();
map.put("spring.grpc.client.default-channel.idle-timeout", "1000");
map.put("spring.grpc.client.default-channel.keep-alive-time", "60000");
map.put("spring.grpc.client.default-channel.keep-alive-timeout", "5000");
map.put("spring.grpc.client.default-channel.idle-timeout", "1");
map.put("spring.grpc.client.default-channel.keep-alive-time", "60");
map.put("spring.grpc.client.default-channel.keep-alive-timeout", "5");
GrpcClientProperties properties = bindProperties(map);
var defaultChannel = properties.getDefaultChannel();
assertThat(defaultChannel.getIdleTimeout()).isEqualTo(Duration.ofSeconds(1));
@@ -161,4 +172,104 @@ class GrpcClientPropertiesTests {
}
@Nested
class GetChannelAPI {
@Test
void withDefaultNameReturnsDefaultChannel() {
var properties = new GrpcClientProperties();
var defaultChannel = properties.getChannel("default");
assertThat(properties).extracting("defaultChannel").isSameAs(defaultChannel);
assertThat(properties).extracting("channels", InstanceOfAssertFactories.MAP).isEmpty();
}
@Test
void withKnownNameReturnsKnownChannel() {
Map<String, String> map = new HashMap<>();
// we have to at least bind one property or bind() fails
map.put("spring.grpc.client.channels.c1.enable-keep-alive", "false");
GrpcClientProperties properties = bindProperties(map);
var channel = properties.getChannel("c1");
assertThat(properties).extracting("channels", InstanceOfAssertFactories.MAP)
.containsExactly(entry("c1", channel));
}
@Test
void withUnknownNameReturnsNewChannelWithCopiedDefaults() {
var defaultChannel = new ChannelConfig();
defaultChannel.setAddress("static://my-server:9999");
defaultChannel.setDefaultLoadBalancingPolicy("custom");
defaultChannel.getHealth().setEnabled(true);
defaultChannel.getHealth().setServiceName("custom-service");
defaultChannel.setEnableKeepAlive(true);
defaultChannel.setIdleTimeout(Duration.ofMinutes(1));
defaultChannel.setKeepAliveTime(Duration.ofMinutes(4));
defaultChannel.setKeepAliveTimeout(Duration.ofMinutes(6));
defaultChannel.setKeepAliveWithoutCalls(true);
defaultChannel.setMaxInboundMessageSize(DataSize.ofMegabytes(100));
defaultChannel.setMaxInboundMetadataSize(DataSize.ofMegabytes(200));
defaultChannel.setUserAgent("me");
defaultChannel.getSsl().setEnabled(true);
defaultChannel.getSsl().setBundle("custom-bundle");
var properties = newProperties(defaultChannel, Map.of());
var newChannel = properties.getChannel("new-channel");
assertThat(newChannel).usingRecursiveComparison().ignoringFields("address").isEqualTo(defaultChannel);
assertThat(newChannel).hasFieldOrPropertyWithValue("address", "static://new-channel");
assertThat(properties).extracting("channels", InstanceOfAssertFactories.MAP).isEmpty();
}
}
@Nested
class GetTargetAPI {
@Test
void channelWithStaticAddressReturnsStrippedAddress() {
var defaultChannel = new ChannelConfig();
var channel1 = new ChannelConfig();
channel1.setAddress("static://my-server:8888");
var properties = newProperties(defaultChannel, Map.of("c1", channel1));
assertThat(properties.getTarget("c1")).isEqualTo("my-server:8888");
assertThat(properties).extracting("channels", InstanceOfAssertFactories.MAP)
.containsExactly(entry("c1", channel1));
}
@Test
void channelWithTcpAddressReturnsStrippedAddress() {
var defaultChannel = new ChannelConfig();
var channel1 = new ChannelConfig();
channel1.setAddress("tcp://my-server:8888");
var properties = newProperties(defaultChannel, Map.of("c1", channel1));
assertThat(properties.getTarget("c1")).isEqualTo("my-server:8888");
assertThat(properties).extracting("channels", InstanceOfAssertFactories.MAP)
.containsExactly(entry("c1", channel1));
}
@Test
void channelWithAddressPropertyPlaceholdersPopulatesFromEnvironment() {
var defaultChannel = new ChannelConfig();
var channel1 = new ChannelConfig();
channel1.setAddress("my-server-${channelName}:8888");
var properties = newProperties(defaultChannel, Map.of("c1", channel1));
var env = new MockEnvironment();
env.setProperty("channelName", "foo");
properties.setEnvironment(env);
assertThat(properties.getTarget("c1")).isEqualTo("my-server-foo:8888");
}
}
@Nested
class CopyDefaultsAPI {
@Test
void copyFromDefaultChannel() {
var properties = new GrpcClientProperties();
var defaultChannel = properties.getDefaultChannel();
var newChannel = defaultChannel.copy();
assertThat(newChannel).usingRecursiveComparison().isEqualTo(defaultChannel);
}
}
}