diff --git a/spring-boot-project/spring-boot-autoconfigure/src/main/java/org/springframework/boot/autoconfigure/rsocket/RSocketProperties.java b/spring-boot-project/spring-boot-autoconfigure/src/main/java/org/springframework/boot/autoconfigure/rsocket/RSocketProperties.java index cf18235260..71e5ce4785 100644 --- a/spring-boot-project/spring-boot-autoconfigure/src/main/java/org/springframework/boot/autoconfigure/rsocket/RSocketProperties.java +++ b/spring-boot-project/spring-boot-autoconfigure/src/main/java/org/springframework/boot/autoconfigure/rsocket/RSocketProperties.java @@ -1,5 +1,5 @@ /* - * Copyright 2012-2019 the original author or authors. + * Copyright 2012-2020 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. diff --git a/spring-boot-project/spring-boot/src/main/java/org/springframework/boot/rsocket/netty/NettyRSocketServerFactory.java b/spring-boot-project/spring-boot/src/main/java/org/springframework/boot/rsocket/netty/NettyRSocketServerFactory.java index 5859fd50fb..278d066e41 100644 --- a/spring-boot-project/spring-boot/src/main/java/org/springframework/boot/rsocket/netty/NettyRSocketServerFactory.java +++ b/spring-boot-project/spring-boot/src/main/java/org/springframework/boot/rsocket/netty/NettyRSocketServerFactory.java @@ -151,41 +151,27 @@ public class NettyRSocketServerFactory implements RSocketServerFactory, Configur } private ServerTransport createWebSocketTransport() { - HttpServer httpServer; + HttpServer httpServer = HttpServer.create(); if (this.resourceFactory != null) { - httpServer = HttpServer.create().runOn(this.resourceFactory.getLoopResources()) - .bindAddress(this::getListenAddress); + httpServer = httpServer.runOn(this.resourceFactory.getLoopResources()); } - else { - InetSocketAddress listenAddress = this.getListenAddress(); - httpServer = HttpServer.create().host(listenAddress.getHostName()).port(listenAddress.getPort()); - } - if (this.ssl != null && this.ssl.isEnabled()) { SslServerCustomizer sslServerCustomizer = new SslServerCustomizer(this.ssl, null, this.sslStoreProvider); httpServer = sslServerCustomizer.apply(httpServer); } - - return WebsocketServerTransport.create(httpServer); + return WebsocketServerTransport.create(httpServer.bindAddress(this::getListenAddress)); } private ServerTransport createTcpTransport() { - TcpServer tcpServer; + TcpServer tcpServer = TcpServer.create(); if (this.resourceFactory != null) { - tcpServer = TcpServer.create().runOn(this.resourceFactory.getLoopResources()) - .bindAddress(this::getListenAddress); + tcpServer = tcpServer.runOn(this.resourceFactory.getLoopResources()); } - else { - InetSocketAddress listenAddress = this.getListenAddress(); - tcpServer = TcpServer.create().host(listenAddress.getHostName()).port(listenAddress.getPort()); - } - if (this.ssl != null && this.ssl.isEnabled()) { TcpSslServerCustomizer sslServerCustomizer = new TcpSslServerCustomizer(this.ssl, this.sslStoreProvider); tcpServer = sslServerCustomizer.apply(tcpServer); } - - return TcpServerTransport.create(tcpServer); + return TcpServerTransport.create(tcpServer.bindAddress(this::getListenAddress)); } private InetSocketAddress getListenAddress() { @@ -201,9 +187,6 @@ public class NettyRSocketServerFactory implements RSocketServerFactory, Configur super(ssl, null, sslStoreProvider); } - // This does not override the apply in parent - currently just leveraging the - // parent for its "getContextBuilder()" method. This should be refactored when - // we add the concept of http/tcp customizers for RSocket. private TcpServer apply(TcpServer server) { try { return server.secure((contextSpec) -> contextSpec.sslContext(getContextBuilder())); diff --git a/spring-boot-project/spring-boot/src/main/java/org/springframework/boot/rsocket/server/ConfigurableRSocketServerFactory.java b/spring-boot-project/spring-boot/src/main/java/org/springframework/boot/rsocket/server/ConfigurableRSocketServerFactory.java index afbf549ba2..40de9b3f7f 100644 --- a/spring-boot-project/spring-boot/src/main/java/org/springframework/boot/rsocket/server/ConfigurableRSocketServerFactory.java +++ b/spring-boot-project/spring-boot/src/main/java/org/springframework/boot/rsocket/server/ConfigurableRSocketServerFactory.java @@ -1,5 +1,5 @@ /* - * Copyright 2012-2019 the original author or authors. + * Copyright 2012-2020 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. diff --git a/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/rsocket/netty/NettyRSocketServerFactoryTests.java b/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/rsocket/netty/NettyRSocketServerFactoryTests.java index d329c6dbd4..30cc5bfb67 100644 --- a/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/rsocket/netty/NettyRSocketServerFactoryTests.java +++ b/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/rsocket/netty/NettyRSocketServerFactoryTests.java @@ -18,7 +18,6 @@ package org.springframework.boot.rsocket.netty; import java.net.InetSocketAddress; import java.nio.channels.ClosedChannelException; -import java.time.Duration; import java.util.Arrays; import java.util.concurrent.Callable; @@ -38,12 +37,13 @@ import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; import org.mockito.InOrder; import reactor.core.publisher.Mono; +import reactor.netty.http.client.HttpClient; import reactor.netty.tcp.TcpClient; import reactor.test.StepVerifier; import org.springframework.boot.rsocket.server.RSocketServer; -import org.springframework.boot.rsocket.server.RSocketServerCustomizer; import org.springframework.boot.rsocket.server.RSocketServer.Transport; +import org.springframework.boot.rsocket.server.RSocketServerCustomizer; import org.springframework.boot.web.server.Ssl; import org.springframework.core.codec.CharSequenceEncoder; import org.springframework.core.codec.StringDecoder; @@ -55,7 +55,6 @@ import org.springframework.util.SocketUtils; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatExceptionOfType; -import static org.assertj.core.api.Assertions.assertThatThrownBy; import static org.mockito.ArgumentMatchers.any; import static org.mockito.BDDMockito.will; import static org.mockito.Mockito.inOrder; @@ -74,10 +73,11 @@ class NettyRSocketServerFactoryTests { private RSocketRequester requester; - private static final Duration TIMEOUT = Duration.ofSeconds(3); - @AfterEach void tearDown() { + if (this.requester != null) { + this.requester.rsocketClient().dispose(); + } if (this.server != null) { try { this.server.stop(); @@ -86,9 +86,6 @@ class NettyRSocketServerFactoryTests { // Ignore } } - if (this.requester != null) { - this.requester.rsocketClient().dispose(); - } } private NettyRSocketServerFactory getFactory() { @@ -105,11 +102,9 @@ class NettyRSocketServerFactoryTests { this.server.start(); return port; }); - this.requester = createRSocketTcpClient(false); - String payload = "test payload"; - String response = this.requester.route("test").data(payload).retrieveMono(String.class).block(TIMEOUT); + this.requester = createRSocketTcpClient(); assertThat(this.server.address().getPort()).isEqualTo(specificPort); - assertThat(response).isEqualTo(payload); + checkEchoRequest(); } @Test @@ -118,10 +113,8 @@ class NettyRSocketServerFactoryTests { factory.setTransport(RSocketServer.Transport.WEBSOCKET); this.server = factory.create(new EchoRequestResponseAcceptor()); this.server.start(); - this.requester = createRSocketWebSocketClient(false); - String payload = "test payload"; - String response = this.requester.route("test").data(payload).retrieveMono(String.class).block(TIMEOUT); - assertThat(response).isEqualTo(payload); + this.requester = createRSocketWebSocketClient(); + checkEchoRequest(); } @Test @@ -133,10 +126,8 @@ class NettyRSocketServerFactoryTests { factory.setResourceFactory(resourceFactory); this.server = factory.create(new EchoRequestResponseAcceptor()); this.server.start(); - this.requester = createRSocketWebSocketClient(false); - String payload = "test payload"; - String response = this.requester.route("test").data(payload).retrieveMono(String.class).block(TIMEOUT); - assertThat(response).isEqualTo(payload); + this.requester = createRSocketWebSocketClient(); + checkEchoRequest(); } @Test @@ -176,6 +167,12 @@ class NettyRSocketServerFactoryTests { testBasicSslWithKeyStore("src/test/resources/test.jks", "password", Transport.WEBSOCKET); } + private void checkEchoRequest() { + String payload = "test payload"; + Mono response = this.requester.route("test").data(payload).retrieveMono(String.class); + StepVerifier.create(response).expectNext(payload).verifyComplete(); + } + private void testBasicSslWithKeyStore(String keyStore, String keyPassword, Transport transport) { NettyRSocketServerFactory factory = getFactory(); factory.setTransport(transport); @@ -185,11 +182,9 @@ class NettyRSocketServerFactoryTests { factory.setSsl(ssl); this.server = factory.create(new EchoRequestResponseAcceptor()); this.server.start(); - this.requester = (transport == Transport.TCP) ? createRSocketTcpClient(true) - : createRSocketWebSocketClient(true); - String payload = "test payload"; - Mono responseMono = this.requester.route("test").data(payload).retrieveMono(String.class); - StepVerifier.create(responseMono).expectNext(payload).verifyComplete(); + this.requester = (transport == Transport.TCP) ? createSecureRSocketTcpClient() + : createSecureRSocketWebSocketClient(); + checkEchoRequest(); } @Test @@ -202,48 +197,54 @@ class NettyRSocketServerFactoryTests { factory.setSsl(ssl); this.server = factory.create(new EchoRequestResponseAcceptor()); this.server.start(); - this.requester = createRSocketTcpClient(false); + this.requester = createRSocketTcpClient(); String payload = "test payload"; Mono responseMono = this.requester.route("test").data(payload).retrieveMono(String.class); StepVerifier.create(responseMono) .verifyErrorSatisfies((ex) -> assertThatExceptionOfType(ClosedChannelException.class)); } - @Test - void websocketTransportSslRejectsInsecureClient() { - NettyRSocketServerFactory factory = getFactory(); - factory.setTransport(Transport.WEBSOCKET); - Ssl ssl = new Ssl(); - ssl.setKeyStore("classpath:test.jks"); - ssl.setKeyPassword("password"); - factory.setSsl(ssl); - this.server = factory.create(new EchoRequestResponseAcceptor()); - this.server.start(); - // For WebSocket, the SSL failure results in a hang on the initial connect call - assertThatThrownBy(() -> createRSocketWebSocketClient(false)).isInstanceOf(IllegalStateException.class) - .hasStackTraceContaining("Timeout on blocking read"); + private RSocketRequester createRSocketTcpClient() { + return createRSocketRequesterBuilder().transport(TcpClientTransport.create(createTcpClient())); } - private RSocketRequester createRSocketTcpClient(boolean ssl) { - TcpClient tcpClient = createTcpClient(ssl); - return createRSocketRequesterBuilder().connect(TcpClientTransport.create(tcpClient)).block(TIMEOUT); + private RSocketRequester createRSocketWebSocketClient() { + return createRSocketRequesterBuilder().transport(WebsocketClientTransport.create(createHttpClient(), "/")); } - private RSocketRequester createRSocketWebSocketClient(boolean ssl) { - TcpClient tcpClient = createTcpClient(ssl); - return createRSocketRequesterBuilder().connect(WebsocketClientTransport.create(tcpClient)).block(TIMEOUT); + private RSocketRequester createSecureRSocketTcpClient() { + return createRSocketRequesterBuilder().transport(TcpClientTransport.create(createSecureTcpClient())); } - private TcpClient createTcpClient(boolean ssl) { + private RSocketRequester createSecureRSocketWebSocketClient() { + return createRSocketRequesterBuilder() + .transport(WebsocketClientTransport.create(createSecureHttpClient(), "/")); + } + + private HttpClient createSecureHttpClient() { + HttpClient httpClient = createHttpClient(); + SslContextBuilder builder = SslContextBuilder.forClient().sslProvider(SslProvider.JDK) + .trustManager(InsecureTrustManagerFactory.INSTANCE); + return httpClient.secure((spec) -> spec.sslContext(builder)); + } + + private HttpClient createHttpClient() { Assertions.assertThat(this.server).isNotNull(); InetSocketAddress address = this.server.address(); - TcpClient tcpClient = TcpClient.create().host(address.getHostName()).port(address.getPort()); - if (ssl) { - SslContextBuilder builder = SslContextBuilder.forClient().sslProvider(SslProvider.JDK) - .trustManager(InsecureTrustManagerFactory.INSTANCE); - tcpClient = tcpClient.secure((spec) -> spec.sslContext(builder)); - } - return tcpClient; + return HttpClient.create().host(address.getHostName()).port(address.getPort()); + } + + private TcpClient createSecureTcpClient() { + TcpClient tcpClient = createTcpClient(); + SslContextBuilder builder = SslContextBuilder.forClient().sslProvider(SslProvider.JDK) + .trustManager(InsecureTrustManagerFactory.INSTANCE); + return tcpClient.secure((spec) -> spec.sslContext(builder)); + } + + private TcpClient createTcpClient() { + Assertions.assertThat(this.server).isNotNull(); + InetSocketAddress address = this.server.address(); + return TcpClient.create().host(address.getHostName()).port(address.getPort()); } private RSocketRequester.Builder createRSocketRequesterBuilder() {