Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
110 changes: 67 additions & 43 deletions rsocket-core/src/main/java/io/rsocket/RSocketFactory.java
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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. */
Expand Down Expand Up @@ -93,10 +96,7 @@ default <T extends Closeable> Start<T> transport(ServerTransport<T> transport) {
public static class ClientRSocketFactory implements ClientTransportAcceptor {
private static final String CLIENT_TAG = "client";

private Supplier<Function<RSocket, RSocket>> acceptor =
() -> rSocket -> new AbstractRSocket() {};

private BiFunction<ConnectionSetupPayload, RSocket, RSocket> biAcceptor;
private SocketAcceptor acceptor = (setup, sendingSocket) -> Mono.just(new AbstractRSocket() {});

private Consumer<Throwable> errorConsumer = Throwable::printStackTrace;
private int mtu = 0;
Expand Down Expand Up @@ -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
*
Expand Down Expand Up @@ -268,18 +273,25 @@ public Start<RSocket> transport(Supplier<ClientTransport> transportClient) {
}

public ClientTransportAcceptor acceptor(Function<RSocket, RSocket> acceptor) {
this.acceptor = () -> acceptor;
return StartClient::new;
return acceptor(() -> acceptor);
}

public ClientTransportAcceptor acceptor(Supplier<Function<RSocket, RSocket>> acceptor) {
this.acceptor = acceptor;
return StartClient::new;
return acceptor(
(SocketAcceptor)
(setup, sendingSocket) -> Mono.just(acceptor.get().apply(sendingSocket)));
}

@Deprecated
public ClientTransportAcceptor acceptor(
BiFunction<ConnectionSetupPayload, RSocket, RSocket> 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;
}

Expand Down Expand Up @@ -346,6 +358,8 @@ public Mono<RSocket> start() {
rSocketRequester = new MultiSubscriberRSocket(rSocketRequester);
}

RSocket wrappedRSocketRequester = plugins.applyRequester(rSocketRequester);

ByteBuf setupFrame =
SetupFrameFlyweight.encode(
allocator,
Expand All @@ -357,34 +371,38 @@ public Mono<RSocket> 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);
});
});
}

Expand Down Expand Up @@ -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<>();
Expand Down Expand Up @@ -644,7 +667,8 @@ private Mono<Void> 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)))
Expand Down
19 changes: 10 additions & 9 deletions rsocket-core/src/main/java/io/rsocket/SocketAcceptor.java
Original file line number Diff line number Diff line change
Expand Up @@ -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<RSocket> accept(ConnectionSetupPayload setup, RSocket sendingSocket);
Expand Down
14 changes: 14 additions & 0 deletions rsocket-core/src/main/java/io/rsocket/plugins/PluginRegistry.java
Original file line number Diff line number Diff line change
Expand Up @@ -18,13 +18,15 @@

import io.rsocket.DuplexConnection;
import io.rsocket.RSocket;
import io.rsocket.SocketAcceptor;
import java.util.ArrayList;
import java.util.List;

public class PluginRegistry {
private List<DuplexConnectionInterceptor> connections = new ArrayList<>();
private List<RSocketInterceptor> requesters = new ArrayList<>();
private List<RSocketInterceptor> responders = new ArrayList<>();
private List<SocketAcceptorInterceptor> socketAcceptorInterceptors = new ArrayList<>();

public PluginRegistry() {}

Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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) {
Expand Down
Original file line number Diff line number Diff line change
@@ -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.
*
* <p>This can be used as an alternative to individual requester and responder {@link
* RSocketInterceptor} plugins.
*/
public @FunctionalInterface interface SocketAcceptorInterceptor
extends Function<SocketAcceptor, SocketAcceptor> {}
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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<Payload> requestResponse(Payload payload) {
calledClient = true;
calledRequester = true;
return reactiveSocket.requestResponse(payload);
}
};

serverPlugin =
responderPlugin =
reactiveSocket ->
new RSocketProxy(reactiveSocket) {
@Override
public Mono<Payload> 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;
Expand All @@ -99,7 +118,8 @@ public void startup() {

server =
RSocketFactory.receive()
.addServerPlugin(serverPlugin)
.addResponderPlugin(responderPlugin)
.addSocketAcceptorPlugin(serverAcceptorPlugin)
.addConnectionPlugin(connectionPlugin)
.errorConsumer(
t -> {
Expand Down Expand Up @@ -138,7 +158,8 @@ public Flux<Payload> requestChannel(Publisher<Payload> payloads) {

client =
RSocketFactory.connect()
.addClientPlugin(clientPlugin)
.addRequesterPlugin(requesterPlugin)
.addSocketAcceptorPlugin(clientAcceptorPlugin)
.addConnectionPlugin(connectionPlugin)
.transport(TcpClientTransport.create(server.address()))
.start()
Expand All @@ -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);
}

Expand Down