diff --git a/rsocket-core/src/main/java/io/rsocket/RSocketFactory.java b/rsocket-core/src/main/java/io/rsocket/RSocketFactory.java index 8a3e3ae25..b6c268464 100644 --- a/rsocket-core/src/main/java/io/rsocket/RSocketFactory.java +++ b/rsocket-core/src/main/java/io/rsocket/RSocketFactory.java @@ -31,11 +31,11 @@ import io.rsocket.internal.ClientSetup; import io.rsocket.internal.ServerSetup; import io.rsocket.keepalive.KeepAliveHandler; -import io.rsocket.lease.*; -import io.rsocket.plugins.DuplexConnectionInterceptor; -import io.rsocket.plugins.PluginRegistry; -import io.rsocket.plugins.Plugins; -import io.rsocket.plugins.RSocketInterceptor; +import io.rsocket.lease.LeaseStats; +import io.rsocket.lease.Leases; +import io.rsocket.lease.RequesterLeaseHandler; +import io.rsocket.lease.ResponderLeaseHandler; +import io.rsocket.plugins.*; import io.rsocket.resume.*; import io.rsocket.transport.ClientTransport; import io.rsocket.transport.ServerTransport; @@ -44,7 +44,10 @@ import io.rsocket.util.MultiSubscriberRSocket; import java.time.Duration; import java.util.Objects; -import java.util.function.*; +import java.util.function.BiFunction; +import java.util.function.Consumer; +import java.util.function.Function; +import java.util.function.Supplier; import reactor.core.publisher.Mono; /** Factory for creating RSocket clients and servers. */ @@ -93,10 +96,7 @@ default Start transport(ServerTransport transport) { public static class ClientRSocketFactory implements ClientTransportAcceptor { private static final String CLIENT_TAG = "client"; - private Supplier> acceptor = - () -> rSocket -> new AbstractRSocket() {}; - - private BiFunction biAcceptor; + private SocketAcceptor acceptor = (setup, sendingSocket) -> Mono.just(new AbstractRSocket() {}); private Consumer errorConsumer = Throwable::printStackTrace; private int mtu = 0; @@ -161,6 +161,11 @@ public ClientRSocketFactory addResponderPlugin(RSocketInterceptor interceptor) { return this; } + public ClientRSocketFactory addSocketAcceptorPlugin(SocketAcceptorInterceptor interceptor) { + plugins.addSocketAcceptorPlugin(interceptor); + return this; + } + /** * Deprecated as Keep-Alive is not optional according to spec * @@ -268,18 +273,25 @@ public Start transport(Supplier transportClient) { } public ClientTransportAcceptor acceptor(Function acceptor) { - this.acceptor = () -> acceptor; - return StartClient::new; + return acceptor(() -> acceptor); } public ClientTransportAcceptor acceptor(Supplier> acceptor) { - this.acceptor = acceptor; - return StartClient::new; + return acceptor( + (SocketAcceptor) + (setup, sendingSocket) -> Mono.just(acceptor.get().apply(sendingSocket))); } + @Deprecated public ClientTransportAcceptor acceptor( BiFunction biAcceptor) { - this.biAcceptor = biAcceptor; + return acceptor( + (SocketAcceptor) + (setup, sendingSocket) -> Mono.just(biAcceptor.apply(setup, sendingSocket))); + } + + public ClientTransportAcceptor acceptor(SocketAcceptor acceptor) { + this.acceptor = acceptor; return StartClient::new; } @@ -346,6 +358,8 @@ public Mono start() { rSocketRequester = new MultiSubscriberRSocket(rSocketRequester); } + RSocket wrappedRSocketRequester = plugins.applyRequester(rSocketRequester); + ByteBuf setupFrame = SetupFrameFlyweight.encode( allocator, @@ -357,34 +371,38 @@ public Mono start() { dataMimeType, setupPayload); - RSocket wrappedRSocketRequester = plugins.applyRequester(rSocketRequester); - - RSocket rSocketHandler; - if (biAcceptor != null) { - ConnectionSetupPayload setup = ConnectionSetupPayload.create(setupFrame); - rSocketHandler = biAcceptor.apply(setup, wrappedRSocketRequester); - } else { - rSocketHandler = acceptor.get().apply(wrappedRSocketRequester); - } - - RSocket wrappedRSocketHandler = plugins.applyResponder(rSocketHandler); - - ResponderLeaseHandler responderLeaseHandler = - isLeaseEnabled - ? new ResponderLeaseHandler.Impl<>( - CLIENT_TAG, allocator, leases.sender(), errorConsumer, leases.stats()) - : ResponderLeaseHandler.None; - - RSocket rSocketResponder = - new RSocketResponder( - allocator, - multiplexer.asServerConnection(), - wrappedRSocketHandler, - payloadDecoder, - errorConsumer, - responderLeaseHandler); + ConnectionSetupPayload setup = ConnectionSetupPayload.create(setupFrame); + + return plugins + .applySocketAcceptorInterceptor(acceptor) + .accept(setup, wrappedRSocketRequester) + .flatMap( + rSocketHandler -> { + RSocket wrappedRSocketHandler = plugins.applyResponder(rSocketHandler); + + ResponderLeaseHandler responderLeaseHandler = + isLeaseEnabled + ? new ResponderLeaseHandler.Impl<>( + CLIENT_TAG, + allocator, + leases.sender(), + errorConsumer, + leases.stats()) + : ResponderLeaseHandler.None; + + RSocket rSocketResponder = + new RSocketResponder( + allocator, + multiplexer.asServerConnection(), + wrappedRSocketHandler, + payloadDecoder, + errorConsumer, + responderLeaseHandler); - return wrappedConnection.sendOne(setupFrame).thenReturn(wrappedRSocketRequester); + return wrappedConnection + .sendOne(setupFrame) + .thenReturn(wrappedRSocketRequester); + }); }); } @@ -476,6 +494,11 @@ public ServerRSocketFactory addResponderPlugin(RSocketInterceptor interceptor) { return this; } + public ServerRSocketFactory addSocketAcceptorPlugin(SocketAcceptorInterceptor interceptor) { + plugins.addSocketAcceptorPlugin(interceptor); + return this; + } + public ServerTransportAcceptor acceptor(SocketAcceptor acceptor) { this.acceptor = acceptor; return new ServerStart<>(); @@ -644,7 +667,8 @@ private Mono acceptSetup( } RSocket wrappedRSocketRequester = plugins.applyRequester(rSocketRequester); - return acceptor + return plugins + .applySocketAcceptorInterceptor(acceptor) .accept(setupPayload, wrappedRSocketRequester) .onErrorResume( err -> sendError(multiplexer, rejectedSetupError(err)).then(Mono.error(err))) diff --git a/rsocket-core/src/main/java/io/rsocket/SocketAcceptor.java b/rsocket-core/src/main/java/io/rsocket/SocketAcceptor.java index 0f6b99d0e..85c731eea 100644 --- a/rsocket-core/src/main/java/io/rsocket/SocketAcceptor.java +++ b/rsocket-core/src/main/java/io/rsocket/SocketAcceptor.java @@ -20,20 +20,21 @@ import reactor.core.publisher.Mono; /** - * {@code RSocket} is a full duplex protocol where a client and server are identical in terms of - * both having the capability to initiate requests to their peer. This interface provides the - * contract where a server accepts a new {@code RSocket} for sending requests to the peer and - * returns a new {@code RSocket} that will be used to accept requests from it's peer. + * RSocket is a full duplex protocol where a client and server are identical in terms of both having + * the capability to initiate requests to their peer. This interface provides the contract where a + * client or server handles the {@code setup} for a new connection and creates a responder {@code + * RSocket} for accepting requests from the remote peer. */ public interface SocketAcceptor { /** - * Accepts a new {@code RSocket} used to send requests to the peer and returns another {@code - * RSocket} that is used for accepting requests from the peer. + * Handle the {@code SETUP} frame for a new connection and create a responder {@code RSocket} for + * handling requests from the remote peer. * - * @param setup Setup as sent by the client. - * @param sendingSocket Socket used to send requests to the peer. - * @return Socket to accept requests from the peer. + * @param setup the {@code setup} received from a client in a server scenario, or in a client + * scenario this is the setup about to be sent to the server. + * @param sendingSocket socket for sending requests to the remote peer. + * @return {@code RSocket} to accept requests with. * @throws SetupException If the acceptor needs to reject the setup of this socket. */ Mono accept(ConnectionSetupPayload setup, RSocket sendingSocket); diff --git a/rsocket-core/src/main/java/io/rsocket/plugins/PluginRegistry.java b/rsocket-core/src/main/java/io/rsocket/plugins/PluginRegistry.java index 676cfc19c..e3a19367c 100644 --- a/rsocket-core/src/main/java/io/rsocket/plugins/PluginRegistry.java +++ b/rsocket-core/src/main/java/io/rsocket/plugins/PluginRegistry.java @@ -18,6 +18,7 @@ import io.rsocket.DuplexConnection; import io.rsocket.RSocket; +import io.rsocket.SocketAcceptor; import java.util.ArrayList; import java.util.List; @@ -25,6 +26,7 @@ public class PluginRegistry { private List connections = new ArrayList<>(); private List requesters = new ArrayList<>(); private List responders = new ArrayList<>(); + private List socketAcceptorInterceptors = new ArrayList<>(); public PluginRegistry() {} @@ -58,6 +60,10 @@ public void addResponderPlugin(RSocketInterceptor interceptor) { responders.add(interceptor); } + public void addSocketAcceptorPlugin(SocketAcceptorInterceptor interceptor) { + socketAcceptorInterceptors.add(interceptor); + } + /** Deprecated. Use {@link #applyRequester(RSocket)} instead */ @Deprecated public RSocket applyClient(RSocket rSocket) { @@ -86,6 +92,14 @@ public RSocket applyResponder(RSocket rSocket) { return rSocket; } + public SocketAcceptor applySocketAcceptorInterceptor(SocketAcceptor acceptor) { + for (SocketAcceptorInterceptor i : socketAcceptorInterceptors) { + acceptor = i.apply(acceptor); + } + + return acceptor; + } + public DuplexConnection applyConnection( DuplexConnectionInterceptor.Type type, DuplexConnection connection) { for (DuplexConnectionInterceptor i : connections) { diff --git a/rsocket-core/src/main/java/io/rsocket/plugins/SocketAcceptorInterceptor.java b/rsocket-core/src/main/java/io/rsocket/plugins/SocketAcceptorInterceptor.java new file mode 100644 index 000000000..c9201ca5b --- /dev/null +++ b/rsocket-core/src/main/java/io/rsocket/plugins/SocketAcceptorInterceptor.java @@ -0,0 +1,29 @@ +/* + * Copyright 2002-2019 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 io.rsocket.plugins; + +import io.rsocket.SocketAcceptor; +import java.util.function.Function; + +/** + * Contract to decorate a {@link SocketAcceptor}, providing access to connection {@code setup} + * information and the ability to also decorate the sockets for requesting and responding. + * + *

This can be used as an alternative to individual requester and responder {@link + * RSocketInterceptor} plugins. + */ +public @FunctionalInterface interface SocketAcceptorInterceptor + extends Function {} diff --git a/rsocket-examples/src/test/java/io/rsocket/integration/IntegrationTest.java b/rsocket-examples/src/test/java/io/rsocket/integration/IntegrationTest.java index 627b1d7da..c7dfe34c6 100644 --- a/rsocket-examples/src/test/java/io/rsocket/integration/IntegrationTest.java +++ b/rsocket-examples/src/test/java/io/rsocket/integration/IntegrationTest.java @@ -29,6 +29,7 @@ import io.rsocket.RSocketFactory; import io.rsocket.plugins.DuplexConnectionInterceptor; import io.rsocket.plugins.RSocketInterceptor; +import io.rsocket.plugins.SocketAcceptorInterceptor; import io.rsocket.test.TestSubscriber; import io.rsocket.transport.netty.client.TcpClientTransport; import io.rsocket.transport.netty.server.CloseableChannel; @@ -48,34 +49,52 @@ public class IntegrationTest { - private static final RSocketInterceptor clientPlugin; - private static final RSocketInterceptor serverPlugin; + private static final RSocketInterceptor requesterPlugin; + private static final RSocketInterceptor responderPlugin; + private static final SocketAcceptorInterceptor clientAcceptorPlugin; + private static final SocketAcceptorInterceptor serverAcceptorPlugin; private static final DuplexConnectionInterceptor connectionPlugin; - public static volatile boolean calledClient = false; - public static volatile boolean calledServer = false; + public static volatile boolean calledRequester = false; + public static volatile boolean calledResponder = false; + public static volatile boolean calledClientAcceptor = false; + public static volatile boolean calledServerAcceptor = false; public static volatile boolean calledFrame = false; static { - clientPlugin = + requesterPlugin = reactiveSocket -> new RSocketProxy(reactiveSocket) { @Override public Mono requestResponse(Payload payload) { - calledClient = true; + calledRequester = true; return reactiveSocket.requestResponse(payload); } }; - serverPlugin = + responderPlugin = reactiveSocket -> new RSocketProxy(reactiveSocket) { @Override public Mono requestResponse(Payload payload) { - calledServer = true; + calledResponder = true; return reactiveSocket.requestResponse(payload); } }; + clientAcceptorPlugin = + acceptor -> + (setup, sendingSocket) -> { + calledClientAcceptor = true; + return acceptor.accept(setup, sendingSocket); + }; + + serverAcceptorPlugin = + acceptor -> + (setup, sendingSocket) -> { + calledServerAcceptor = true; + return acceptor.accept(setup, sendingSocket); + }; + connectionPlugin = (type, connection) -> { calledFrame = true; @@ -99,7 +118,8 @@ public void startup() { server = RSocketFactory.receive() - .addServerPlugin(serverPlugin) + .addResponderPlugin(responderPlugin) + .addSocketAcceptorPlugin(serverAcceptorPlugin) .addConnectionPlugin(connectionPlugin) .errorConsumer( t -> { @@ -138,7 +158,8 @@ public Flux requestChannel(Publisher payloads) { client = RSocketFactory.connect() - .addClientPlugin(clientPlugin) + .addRequesterPlugin(requesterPlugin) + .addSocketAcceptorPlugin(clientAcceptorPlugin) .addConnectionPlugin(connectionPlugin) .transport(TcpClientTransport.create(server.address())) .start() @@ -154,8 +175,10 @@ public void teardown() { public void testRequest() { client.requestResponse(DefaultPayload.create("REQUEST", "META")).block(); assertThat("Server did not see the request.", requestCount.get(), is(1)); - assertTrue(calledClient); - assertTrue(calledServer); + assertTrue(calledRequester); + assertTrue(calledResponder); + assertTrue(calledClientAcceptor); + assertTrue(calledServerAcceptor); assertTrue(calledFrame); }