diff --git a/ratis-netty/pom.xml b/ratis-netty/pom.xml index 5688a1fa4a..6b720332c9 100644 --- a/ratis-netty/pom.xml +++ b/ratis-netty/pom.xml @@ -88,6 +88,11 @@ junit-platform-launcher test + + org.mockito + mockito-core + test + diff --git a/ratis-netty/src/main/java/org/apache/ratis/netty/NettyConfigKeys.java b/ratis-netty/src/main/java/org/apache/ratis/netty/NettyConfigKeys.java index fd0906ebfe..e7ba5d8101 100644 --- a/ratis-netty/src/main/java/org/apache/ratis/netty/NettyConfigKeys.java +++ b/ratis-netty/src/main/java/org/apache/ratis/netty/NettyConfigKeys.java @@ -74,6 +74,29 @@ static boolean useEpoll(RaftProperties properties) { static void setUseEpoll(RaftProperties properties, boolean enable) { setBoolean(properties::setBoolean, USE_EPOLL_KEY, enable); } + + String ASYNC_REQUEST_THREAD_POOL_CACHED_KEY = PREFIX + ".async.request.thread.pool.cached"; + // Default to a fixed pool. + // TODO: Refer to https://issues.apache.org/jira/browse/RATIS-2637 + boolean ASYNC_REQUEST_THREAD_POOL_CACHED_DEFAULT = false; + static boolean asyncRequestThreadPoolCached(RaftProperties properties) { + return getBoolean(properties::getBoolean, ASYNC_REQUEST_THREAD_POOL_CACHED_KEY, + ASYNC_REQUEST_THREAD_POOL_CACHED_DEFAULT, getDefaultLog()); + } + static void setAsyncRequestThreadPoolCached(RaftProperties properties, boolean useCached) { + setBoolean(properties::setBoolean, ASYNC_REQUEST_THREAD_POOL_CACHED_KEY, useCached); + } + + String ASYNC_REQUEST_THREAD_POOL_SIZE_KEY = PREFIX + ".async.request.thread.pool.size"; + int ASYNC_REQUEST_THREAD_POOL_SIZE_DEFAULT = 32; + static int asyncRequestThreadPoolSize(RaftProperties properties) { + return getInt(properties::getInt, ASYNC_REQUEST_THREAD_POOL_SIZE_KEY, + ASYNC_REQUEST_THREAD_POOL_SIZE_DEFAULT, getDefaultLog(), + requireMin(0), requireMax(65536)); + } + static void setAsyncRequestThreadPoolSize(RaftProperties properties, int size) { + setInt(properties::setInt, ASYNC_REQUEST_THREAD_POOL_SIZE_KEY, size); + } } interface Client { diff --git a/ratis-netty/src/main/java/org/apache/ratis/netty/server/NettyRpcService.java b/ratis-netty/src/main/java/org/apache/ratis/netty/server/NettyRpcService.java index f7d2805e8a..b93137d196 100644 --- a/ratis-netty/src/main/java/org/apache/ratis/netty/server/NettyRpcService.java +++ b/ratis-netty/src/main/java/org/apache/ratis/netty/server/NettyRpcService.java @@ -21,8 +21,6 @@ import org.apache.ratis.netty.NettyConfigKeys; import org.apache.ratis.netty.NettyRpcProxy; import org.apache.ratis.util.NettyUtils; -import org.apache.ratis.protocol.GroupInfoReply; -import org.apache.ratis.protocol.GroupListReply; import org.apache.ratis.protocol.RaftClientReply; import org.apache.ratis.protocol.RaftPeerId; import org.apache.ratis.rpc.SupportedRpcType; @@ -42,16 +40,20 @@ import org.apache.ratis.proto.netty.NettyProtos.RaftNettyServerReplyProto; import org.apache.ratis.proto.netty.NettyProtos.RaftNettyServerRequestProto; import org.apache.ratis.util.CodeInjectionForTesting; +import org.apache.ratis.util.ConcurrentUtils; import org.apache.ratis.util.JavaUtils; import org.apache.ratis.util.MemoizedSupplier; import org.apache.ratis.util.ProtoUtils; +import org.apache.ratis.util.TimeDuration; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import java.io.IOException; import java.net.InetSocketAddress; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CompletionException; +import java.util.concurrent.ExecutorService; import java.util.concurrent.TimeUnit; -import java.util.Objects; /** * A netty server endpoint that acts as the communication layer. @@ -87,12 +89,43 @@ public static Builder newBuilder() { private final MemoizedSupplier channel; private final InetSocketAddress socketAddress; - @ChannelHandler.Sharable + private final ExecutorService requestExecutor; + class InboundHandler extends SimpleChannelInboundHandler { + /** + * Tail of this channel's chain of request-handling tasks. + * Requests on a channel must be handled in arrival order. + */ + private CompletableFuture tail = CompletableFuture.completedFuture(null); + @Override protected void channelRead0(ChannelHandlerContext ctx, RaftNettyServerRequestProto proto) { - final RaftNettyServerReplyProto reply = handle(proto); - ctx.writeAndFlush(reply); + tail = tail.handleAsync((prev, prevError) -> { + final CompletableFuture replyFuture; + try { + replyFuture = handleAsync(proto); + } catch (Exception e) { + // No request context to build a reply; fail fast by closing the channel. + LOG.warn("{}: Failed to handle request; closing the channel.", getId(), e); + ctx.close(); + return null; + } + replyFuture.whenComplete((reply, e) -> { + if (e != null) { + LOG.warn("{}: Failed to handle request; closing the channel.", getId(), e); + ctx.close(); + } else { + ctx.writeAndFlush(reply); + } + }); + return null; + }, requestExecutor); + } + + @Override + public void exceptionCaught(ChannelHandlerContext ctx, Throwable cause) { + LOG.warn("{}: exceptionCaught on channel {}; closing it.", getId(), ctx.channel(), cause); + ctx.close(); } } @@ -130,6 +163,11 @@ protected void initChannel(SocketChannel ch) { .handler(new LoggingHandler(LogLevel.INFO)) .childHandler(initializer) .bind(socketAddress)); + + this.requestExecutor = ConcurrentUtils.newThreadPoolWithMax( + NettyConfigKeys.Server.asyncRequestThreadPoolCached(server.getProperties()), + NettyConfigKeys.Server.asyncRequestThreadPoolSize(server.getProperties()), + server.getId() + "-request-"); } @Override @@ -155,6 +193,8 @@ public void startImpl() throws IOException { @Override public void closeImpl() throws IOException { + ConcurrentUtils.shutdownAndWait(TimeDuration.ONE_SECOND, requestExecutor, + timeout -> LOG.warn("{}: requestExecutor shutdown timeout in {}", this, timeout)); final ChannelFuture f = getChannel().close(); f.syncUninterruptibly(); bossGroup.shutdownGracefully(0, 100, TimeUnit.MILLISECONDS); @@ -181,113 +221,111 @@ public InetSocketAddress getInetSocketAddress() { } } - RaftNettyServerReplyProto handle(RaftNettyServerRequestProto proto) { + CompletableFuture handleAsync(RaftNettyServerRequestProto proto) { RaftRpcRequestProto rpcRequest = null; try { + final CompletableFuture replyFuture; switch (proto.getRaftNettyServerRequestCase()) { case REQUESTVOTEREQUEST: - final RequestVoteRequestProto request = proto.getRequestVoteRequest(); - rpcRequest = request.getServerRequest(); - final RequestVoteReplyProto reply = server.requestVote(request); - return RaftNettyServerReplyProto.newBuilder() - .setRequestVoteReply(reply) - .build(); + // requestVote has no async variant; it is fast and does not block on commit. + final RequestVoteRequestProto requestVoteRequest = proto.getRequestVoteRequest(); + rpcRequest = requestVoteRequest.getServerRequest(); + replyFuture = CompletableFuture.completedFuture(RaftNettyServerReplyProto.newBuilder() + .setRequestVoteReply(server.requestVote(requestVoteRequest)) + .build()); + break; case TRANSFERLEADERSHIPREQUEST: final TransferLeadershipRequestProto transferLeadershipRequest = proto.getTransferLeadershipRequest(); rpcRequest = transferLeadershipRequest.getRpcRequest(); - final RaftClientReply transferLeadershipReply = server.transferLeadership( - ClientProtoUtils.toTransferLeadershipRequest(transferLeadershipRequest)); - return RaftNettyServerReplyProto.newBuilder() - .setRaftClientReply(ClientProtoUtils.toRaftClientReplyProto(transferLeadershipReply)) - .build(); + replyFuture = server.transferLeadershipAsync( + ClientProtoUtils.toTransferLeadershipRequest(transferLeadershipRequest)) + .thenApply(NettyRpcService::toRaftClientReply); + break; case STARTLEADERELECTIONREQUEST: + // startLeaderElection has no async variant; it is fast and does not block on commit. final StartLeaderElectionRequestProto startLeaderElectionRequest = proto.getStartLeaderElectionRequest(); rpcRequest = startLeaderElectionRequest.getServerRequest(); - final StartLeaderElectionReplyProto startLeaderElectionReply = - server.startLeaderElection(startLeaderElectionRequest); - return RaftNettyServerReplyProto.newBuilder().setStartLeaderElectionReply(startLeaderElectionReply).build(); + replyFuture = CompletableFuture.completedFuture(RaftNettyServerReplyProto.newBuilder() + .setStartLeaderElectionReply(server.startLeaderElection(startLeaderElectionRequest)) + .build()); + break; case SNAPSHOTMANAGEMENTREQUEST: final SnapshotManagementRequestProto snapshotManagementRequest = proto.getSnapshotManagementRequest(); rpcRequest = snapshotManagementRequest.getRpcRequest(); - final RaftClientReply snapshotManagementReply = server.snapshotManagement( - ClientProtoUtils.toSnapshotManagementRequest(snapshotManagementRequest)); - return RaftNettyServerReplyProto.newBuilder() - .setRaftClientReply(ClientProtoUtils.toRaftClientReplyProto(snapshotManagementReply)) - .build(); + replyFuture = server.snapshotManagementAsync( + ClientProtoUtils.toSnapshotManagementRequest(snapshotManagementRequest)) + .thenApply(NettyRpcService::toRaftClientReply); + break; case LEADERELECTIONMANAGEMENTREQUEST: final LeaderElectionManagementRequestProto leaderElectionManagementRequest = proto.getLeaderElectionManagementRequest(); rpcRequest = leaderElectionManagementRequest.getRpcRequest(); - final RaftClientReply leaderElectionManagementReply = server.leaderElectionManagement( - ClientProtoUtils.toLeaderElectionManagementRequest(leaderElectionManagementRequest)); - return RaftNettyServerReplyProto.newBuilder() - .setRaftClientReply(ClientProtoUtils.toRaftClientReplyProto(leaderElectionManagementReply)) - .build(); + replyFuture = server.leaderElectionManagementAsync( + ClientProtoUtils.toLeaderElectionManagementRequest(leaderElectionManagementRequest)) + .thenApply(NettyRpcService::toRaftClientReply); + break; case APPENDENTRIESREQUEST: final AppendEntriesRequestProto appendEntriesRequest = proto.getAppendEntriesRequest(); rpcRequest = appendEntriesRequest.getServerRequest(); - final AppendEntriesReplyProto appendEntriesReply = server.appendEntries(appendEntriesRequest); - return RaftNettyServerReplyProto.newBuilder() - .setAppendEntriesReply(appendEntriesReply) - .build(); + replyFuture = server.appendEntriesAsync(appendEntriesRequest) + .thenApply(reply -> RaftNettyServerReplyProto.newBuilder() + .setAppendEntriesReply(reply) + .build()); + break; case INSTALLSNAPSHOTREQUEST: + // installSnapshot has no async variant; it runs on this per-channel serialized path. final InstallSnapshotRequestProto installSnapshotRequest = proto.getInstallSnapshotRequest(); rpcRequest = installSnapshotRequest.getServerRequest(); - final InstallSnapshotReplyProto installSnapshotReply = server.installSnapshot(installSnapshotRequest); - return RaftNettyServerReplyProto.newBuilder() - .setInstallSnapshotReply(installSnapshotReply) - .build(); + replyFuture = CompletableFuture.completedFuture(RaftNettyServerReplyProto.newBuilder() + .setInstallSnapshotReply(server.installSnapshot(installSnapshotRequest)) + .build()); + break; case RAFTCLIENTREQUEST: final RaftClientRequestProto raftClientRequest = proto.getRaftClientRequest(); rpcRequest = raftClientRequest.getRpcRequest(); - final RaftClientReply raftClientReply = server.submitClientRequest( - ClientProtoUtils.toRaftClientRequest(raftClientRequest)); - return RaftNettyServerReplyProto.newBuilder() - .setRaftClientReply(ClientProtoUtils.toRaftClientReplyProto(raftClientReply)) - .build(); + replyFuture = server.submitClientRequestAsync(ClientProtoUtils.toRaftClientRequest(raftClientRequest)) + .thenApply(NettyRpcService::toRaftClientReply); + break; case SETCONFIGURATIONREQUEST: - final SetConfigurationRequestProto configurationRequest = proto.getSetConfigurationRequest(); - rpcRequest = configurationRequest.getRpcRequest(); - final RaftClientReply configurationReply = server.setConfiguration( - ClientProtoUtils.toSetConfigurationRequest(configurationRequest)); - return RaftNettyServerReplyProto.newBuilder() - .setRaftClientReply(ClientProtoUtils.toRaftClientReplyProto(configurationReply)) - .build(); + final SetConfigurationRequestProto setConfigurationRequest = proto.getSetConfigurationRequest(); + rpcRequest = setConfigurationRequest.getRpcRequest(); + replyFuture = server.setConfigurationAsync( + ClientProtoUtils.toSetConfigurationRequest(setConfigurationRequest)) + .thenApply(NettyRpcService::toRaftClientReply); + break; case GROUPMANAGEMENTREQUEST: final GroupManagementRequestProto groupManagementRequest = proto.getGroupManagementRequest(); rpcRequest = groupManagementRequest.getRpcRequest(); - final RaftClientReply groupManagementReply = server.groupManagement( - ClientProtoUtils.toGroupManagementRequest(groupManagementRequest)); - return RaftNettyServerReplyProto.newBuilder() - .setRaftClientReply(ClientProtoUtils.toRaftClientReplyProto(groupManagementReply)) - .build(); + replyFuture = server.groupManagementAsync(ClientProtoUtils.toGroupManagementRequest(groupManagementRequest)) + .thenApply(NettyRpcService::toRaftClientReply); + break; case GROUPLISTREQUEST: final GroupListRequestProto groupListRequest = proto.getGroupListRequest(); rpcRequest = groupListRequest.getRpcRequest(); - final GroupListReply groupListReply = server.getGroupList( - ClientProtoUtils.toGroupListRequest(groupListRequest)); - return RaftNettyServerReplyProto.newBuilder() - .setGroupListReply(ClientProtoUtils.toGroupListReplyProto(groupListReply)) - .build(); + replyFuture = server.getGroupListAsync(ClientProtoUtils.toGroupListRequest(groupListRequest)) + .thenApply(reply -> RaftNettyServerReplyProto.newBuilder() + .setGroupListReply(ClientProtoUtils.toGroupListReplyProto(reply)) + .build()); + break; case GROUPINFOREQUEST: final GroupInfoRequestProto groupInfoRequest = proto.getGroupInfoRequest(); rpcRequest = groupInfoRequest.getRpcRequest(); - final GroupInfoReply groupInfoReply = server.getGroupInfo( - ClientProtoUtils.toGroupInfoRequest(groupInfoRequest)); - return RaftNettyServerReplyProto.newBuilder() - .setGroupInfoReply(ClientProtoUtils.toGroupInfoReplyProto(groupInfoReply)) - .build(); + replyFuture = server.getGroupInfoAsync(ClientProtoUtils.toGroupInfoRequest(groupInfoRequest)) + .thenApply(reply -> RaftNettyServerReplyProto.newBuilder() + .setGroupInfoReply(ClientProtoUtils.toGroupInfoReplyProto(reply)) + .build()); + break; case RAFTNETTYSERVERREQUEST_NOT_SET: throw new IllegalArgumentException("Request case not set in proto: " @@ -296,12 +334,31 @@ RaftNettyServerReplyProto handle(RaftNettyServerRequestProto proto) { throw new UnsupportedOperationException("Request case not supported: " + proto.getRaftNettyServerRequestCase()); } - } catch (IOException ioe) { - return toRaftNettyServerReplyProto( - Objects.requireNonNull(rpcRequest, "rpcRequest = null"), ioe); + + final RaftRpcRequestProto request = rpcRequest; + // Convert an asynchronous failure into an error reply (the client casts it to IOException). + return replyFuture.exceptionally(e -> toRaftNettyServerReplyProto(request, toIOException(e))); + } catch (IOException | RuntimeException e) { + // A synchronous failure before the reply future was created. + if (rpcRequest == null) { + // No request context to build a targeted reply; let InboundHandler close the channel. + throw new IllegalStateException(getId() + ": Failed to handle request " + proto, e); + } + return CompletableFuture.completedFuture(toRaftNettyServerReplyProto(rpcRequest, toIOException(e))); } } + private static RaftNettyServerReplyProto toRaftClientReply(RaftClientReply reply) { + return RaftNettyServerReplyProto.newBuilder() + .setRaftClientReply(ClientProtoUtils.toRaftClientReplyProto(reply)) + .build(); + } + + private static IOException toIOException(Throwable t) { + final Throwable cause = t instanceof CompletionException && t.getCause() != null ? t.getCause() : t; + return cause instanceof IOException ? (IOException) cause : new IOException(cause); + } + private static RaftNettyServerReplyProto toRaftNettyServerReplyProto( RaftRpcRequestProto request, IOException e) { final RaftRpcReplyProto.Builder rpcReply = RaftRpcReplyProto.newBuilder() diff --git a/ratis-netty/src/test/java/org/apache/ratis/netty/server/TestNettyRpcService.java b/ratis-netty/src/test/java/org/apache/ratis/netty/server/TestNettyRpcService.java new file mode 100644 index 0000000000..8579909135 --- /dev/null +++ b/ratis-netty/src/test/java/org/apache/ratis/netty/server/TestNettyRpcService.java @@ -0,0 +1,105 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you 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 + * + * http://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.apache.ratis.netty.server; + +import org.apache.ratis.conf.RaftProperties; +import org.apache.ratis.proto.RaftProtos.RaftRpcRequestProto; +import org.apache.ratis.proto.RaftProtos.RequestVoteRequestProto; +import org.apache.ratis.proto.netty.NettyProtos.RaftNettyServerReplyProto; +import org.apache.ratis.proto.netty.NettyProtos.RaftNettyServerReplyProto.RaftNettyServerReplyCase; +import org.apache.ratis.proto.netty.NettyProtos.RaftNettyServerRequestProto; +import org.apache.ratis.protocol.RaftPeerId; +import org.apache.ratis.server.RaftServer; +import org.apache.ratis.thirdparty.io.netty.channel.ChannelHandlerContext; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; +import org.mockito.Mockito; + +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.TimeUnit; + +/** Tests for {@link NettyRpcService} request handling. */ +public class TestNettyRpcService { + private static final RaftPeerId ID = RaftPeerId.valueOf("s0"); + + private static RaftServer newMockServer() { + final RaftServer server = Mockito.mock(RaftServer.class); + Mockito.when(server.getId()).thenReturn(ID); + Mockito.when(server.getProperties()).thenReturn(new RaftProperties()); + return server; + } + + private static RaftNettyServerRequestProto newRequestVoteProto() { + final RaftRpcRequestProto rpc = RaftRpcRequestProto.newBuilder() + .setRequestorId(ID.toByteString()) + .setReplyId(ID.toByteString()) + .setCallId(1) + .build(); + final RequestVoteRequestProto request = RequestVoteRequestProto.newBuilder() + .setServerRequest(rpc) + .build(); + return RaftNettyServerRequestProto.newBuilder() + .setRequestVoteRequest(request) + .build(); + } + + /** + * A non-{@link java.io.IOException} thrown by the server must be turned into an error reply + * instead of escaping the handler and leaving the client to block until its request timeout. + */ + @Test + public void testHandleReturnsErrorReplyOnRuntimeException() throws Exception { + final RaftServer server = newMockServer(); + Mockito.when(server.requestVote(Mockito.any())).thenThrow(new RuntimeException("injected")); + + final NettyRpcService service = NettyRpcService.newBuilder().setServer(server).build(); + service.start(); + try { + final RaftNettyServerReplyProto reply = service.handleAsync(newRequestVoteProto()).join(); + Assertions.assertEquals(RaftNettyServerReplyCase.EXCEPTIONREPLY, reply.getRaftNettyServerReplyCase()); + } finally { + service.close(); + } + } + + /** Requests must be handled off the Netty I/O event loop, on the request executor thread. */ + @Test + public void testRequestHandledOffEventLoop() throws Exception { + final RaftServer server = newMockServer(); + final CompletableFuture handlingThreadName = new CompletableFuture<>(); + Mockito.when(server.requestVote(Mockito.any())).thenAnswer(invocation -> { + handlingThreadName.complete(Thread.currentThread().getName()); + throw new RuntimeException("injected"); + }); + + final NettyRpcService service = NettyRpcService.newBuilder().setServer(server).build(); + service.start(); + try { + final ChannelHandlerContext ctx = Mockito.mock(ChannelHandlerContext.class); + service.new InboundHandler().channelRead0(ctx, newRequestVoteProto()); + + final String threadName = handlingThreadName.get(5, TimeUnit.SECONDS); + Assertions.assertTrue(threadName.startsWith(ID + "-request-"), + "Request was handled on an unexpected thread: " + threadName); + Assertions.assertNotEquals(Thread.currentThread().getName(), threadName, + "Request was handled on the calling thread, not offloaded"); + } finally { + service.close(); + } + } +}