diff --git a/samples/grpc-server/src/test/java/org/springframework/grpc/sample/GrpcServerIntegrationTests.java b/samples/grpc-server/src/test/java/org/springframework/grpc/sample/GrpcServerIntegrationTests.java index 58176ee..b61383e 100644 --- a/samples/grpc-server/src/test/java/org/springframework/grpc/sample/GrpcServerIntegrationTests.java +++ b/samples/grpc-server/src/test/java/org/springframework/grpc/sample/GrpcServerIntegrationTests.java @@ -18,6 +18,7 @@ package org.springframework.grpc.sample; import static org.assertj.core.api.Assertions.assertThat; +import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledOnOs; @@ -119,6 +120,23 @@ class GrpcServerIntegrationTests { } + @Nested + @SpringBootTest(properties = { "spring.grpc.server.port=0", "spring.grpc.server.ssl.client-auth=REQUIRE", + "spring.grpc.client.channels.test-channel.address=static://0.0.0.0:${local.grpc.port}", + "spring.grpc.client.channels.test-channel.negotiation-type=TLS", + "spring.grpc.client.channels.test-channel.secure=false" }) + @ActiveProfiles("ssl") + @DirtiesContext + @Disabled("Requires client certificate") + class ServerWithClientAuth { + + @Test + void clientChannelWithSsl(@Autowired GrpcChannelFactory channels) { + assertThatResponseIsServedToChannel(channels.createChannel("test-channel").build()); + } + + } + private void assertThatResponseIsServedToChannel(ManagedChannel clientChannel) { SimpleGrpc.SimpleBlockingStub client = SimpleGrpc.newBlockingStub(clientChannel); HelloReply response = client.sayHello(HelloRequest.newBuilder().setName("Alien").build()); diff --git a/spring-grpc-core/src/main/java/org/springframework/grpc/server/DefaultGrpcServerFactory.java b/spring-grpc-core/src/main/java/org/springframework/grpc/server/DefaultGrpcServerFactory.java index 1a91ef0..a67e04f 100644 --- a/spring-grpc-core/src/main/java/org/springframework/grpc/server/DefaultGrpcServerFactory.java +++ b/spring-grpc-core/src/main/java/org/springframework/grpc/server/DefaultGrpcServerFactory.java @@ -21,6 +21,13 @@ import java.util.List; import java.util.Objects; import java.util.Set; +import javax.net.ssl.KeyManagerFactory; +import javax.net.ssl.TrustManagerFactory; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.springframework.grpc.internal.GrpcUtils; + import com.google.common.collect.Lists; import io.grpc.Grpc; @@ -30,10 +37,9 @@ import io.grpc.ServerBuilder; import io.grpc.ServerCredentials; import io.grpc.ServerProvider; import io.grpc.ServerServiceDefinition; - -import org.apache.commons.logging.Log; -import org.apache.commons.logging.LogFactory; -import org.springframework.grpc.internal.GrpcUtils; +import io.grpc.TlsServerCredentials; +import io.grpc.TlsServerCredentials.Builder; +import io.grpc.TlsServerCredentials.ClientAuth; /** * Default implementation for {@link GrpcServerFactory gRPC service factories}. @@ -56,11 +62,25 @@ public class DefaultGrpcServerFactory> implements Grp private final List> serverBuilderCustomizers; + private KeyManagerFactory keyManager; + + private TrustManagerFactory trustManager; + + private ClientAuth clientAuth; + public DefaultGrpcServerFactory(String address, List> serverBuilderCustomizers) { this.address = address; this.serverBuilderCustomizers = Objects.requireNonNull(serverBuilderCustomizers, "serverBuilderCustomizers"); } + public DefaultGrpcServerFactory(String address, List> serverBuilderCustomizers, + KeyManagerFactory keyManager, TrustManagerFactory trustManager, ClientAuth clientAuth) { + this(address, serverBuilderCustomizers); + this.keyManager = keyManager; + this.trustManager = trustManager; + this.clientAuth = clientAuth; + } + protected String address() { return this.address; } @@ -99,7 +119,17 @@ public class DefaultGrpcServerFactory> implements Grp * @return some server credentials (default is insecure) */ protected ServerCredentials credentials() { - return InsecureServerCredentials.create(); + if (this.keyManager == null || port() == -1) { + return InsecureServerCredentials.create(); + } + Builder builder = TlsServerCredentials.newBuilder().keyManager(this.keyManager.getKeyManagers()); + if (this.trustManager != null) { + builder.trustManager(this.trustManager.getTrustManagers()); + } + if (this.clientAuth != null) { + builder.clientAuth(this.clientAuth); + } + return builder.build(); } /** diff --git a/spring-grpc-core/src/main/java/org/springframework/grpc/server/NettyGrpcServerFactory.java b/spring-grpc-core/src/main/java/org/springframework/grpc/server/NettyGrpcServerFactory.java index b5ca59d..50db950 100644 --- a/spring-grpc-core/src/main/java/org/springframework/grpc/server/NettyGrpcServerFactory.java +++ b/spring-grpc-core/src/main/java/org/springframework/grpc/server/NettyGrpcServerFactory.java @@ -19,9 +19,12 @@ package org.springframework.grpc.server; import java.util.List; import javax.net.ssl.KeyManagerFactory; +import javax.net.ssl.TrustManagerFactory; import io.grpc.ServerCredentials; import io.grpc.TlsServerCredentials; +import io.grpc.TlsServerCredentials.Builder; +import io.grpc.TlsServerCredentials.ClientAuth; import io.grpc.netty.NettyServerBuilder; import io.netty.channel.epoll.EpollEventLoopGroup; import io.netty.channel.epoll.EpollServerDomainSocketChannel; @@ -35,12 +38,10 @@ import io.netty.channel.unix.DomainSocketAddress; */ public class NettyGrpcServerFactory extends DefaultGrpcServerFactory { - private KeyManagerFactory keyManager; - - public NettyGrpcServerFactory(String address, KeyManagerFactory keyManager, - List> serverBuilderCustomizers) { - super(address, serverBuilderCustomizers); - this.keyManager = keyManager; + public NettyGrpcServerFactory(String address, + List> serverBuilderCustomizers, KeyManagerFactory keyManager, + TrustManagerFactory trustManager, ClientAuth clientAuth) { + super(address, serverBuilderCustomizers, keyManager, trustManager, clientAuth); } @Override @@ -56,12 +57,4 @@ public class NettyGrpcServerFactory extends DefaultGrpcServerFactory { - private KeyManagerFactory keyManager; - - public ShadedNettyGrpcServerFactory(String address, KeyManagerFactory keyManager, - List> serverBuilderCustomizers) { - super(address, serverBuilderCustomizers); - this.keyManager = keyManager; + public ShadedNettyGrpcServerFactory(String address, + List> serverBuilderCustomizers, KeyManagerFactory keyManager, + TrustManagerFactory trustManager, ClientAuth clientAuth) { + super(address, serverBuilderCustomizers, keyManager, trustManager, clientAuth); } @Override @@ -56,12 +54,4 @@ public class ShadedNettyGrpcServerFactory extends DefaultGrpcServerFactory> builderCustomizers = List .of(mapper::customizeServerBuilder, serverBuilderCustomizers::customize); KeyManagerFactory keyManager = null; + TrustManagerFactory trustManager = null; if (properties.getSsl().isEnabled()) { SslBundle bundle = bundles.getBundle(properties.getSsl().getBundle()); keyManager = bundle.getManagers().getKeyManagerFactory(); + trustManager = bundle.getManagers().getTrustManagerFactory(); } - ShadedNettyGrpcServerFactory factory = new ShadedNettyGrpcServerFactory(properties.getAddress(), keyManager, - builderCustomizers); + ShadedNettyGrpcServerFactory factory = new ShadedNettyGrpcServerFactory(properties.getAddress(), + builderCustomizers, keyManager, trustManager, properties.getSsl().getClientAuth()); grpcServicesDiscoverer.findServices().forEach(factory::addService); return factory; } @@ -82,12 +85,14 @@ class GrpcServerFactoryConfigurations { List> builderCustomizers = List .of(mapper::customizeServerBuilder, serverBuilderCustomizers::customize); KeyManagerFactory keyManager = null; + TrustManagerFactory trustManager = null; if (properties.getSsl().isEnabled()) { SslBundle bundle = bundles.getBundle(properties.getSsl().getBundle()); keyManager = bundle.getManagers().getKeyManagerFactory(); + trustManager = bundle.getManagers().getTrustManagerFactory(); } - NettyGrpcServerFactory factory = new NettyGrpcServerFactory(properties.getAddress(), keyManager, - builderCustomizers); + NettyGrpcServerFactory factory = new NettyGrpcServerFactory(properties.getAddress(), builderCustomizers, + keyManager, trustManager, properties.getSsl().getClientAuth()); grpcServicesDiscoverer.findServices().forEach(factory::addService); return factory; } diff --git a/spring-grpc-spring-boot-autoconfigure/src/main/java/org/springframework/grpc/autoconfigure/server/GrpcServerProperties.java b/spring-grpc-spring-boot-autoconfigure/src/main/java/org/springframework/grpc/autoconfigure/server/GrpcServerProperties.java index 95acf0f..601950b 100644 --- a/spring-grpc-spring-boot-autoconfigure/src/main/java/org/springframework/grpc/autoconfigure/server/GrpcServerProperties.java +++ b/spring-grpc-spring-boot-autoconfigure/src/main/java/org/springframework/grpc/autoconfigure/server/GrpcServerProperties.java @@ -25,6 +25,8 @@ import org.springframework.grpc.internal.GrpcUtils; import org.springframework.util.unit.DataSize; import org.springframework.util.unit.DataUnit; +import io.grpc.TlsServerCredentials.ClientAuth; + @ConfigurationProperties(prefix = "spring.grpc.server") public class GrpcServerProperties { @@ -253,6 +255,11 @@ public class GrpcServerProperties { */ private Boolean enabled; + /** + * Client authentication mode. + */ + private ClientAuth clientAuth = ClientAuth.NONE; + /** * SSL bundle name. */ @@ -284,6 +291,14 @@ public class GrpcServerProperties { this.bundle = bundle; } + public void setClientAuth(ClientAuth clientAuth) { + this.clientAuth = clientAuth; + } + + public ClientAuth getClientAuth() { + return clientAuth; + } + } }