diff --git a/fe/fe-core/src/main/java/org/apache/doris/common/ThriftServer.java b/fe/fe-core/src/main/java/org/apache/doris/common/ThriftServer.java index cdc4bba71d3bb9..898ad23ef599c3 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/common/ThriftServer.java +++ b/fe/fe-core/src/main/java/org/apache/doris/common/ThriftServer.java @@ -33,8 +33,10 @@ import org.apache.thrift.transport.TNonblockingServerSocket; import org.apache.thrift.transport.TNonblockingSocket; import org.apache.thrift.transport.TServerSocket; +import org.apache.thrift.transport.TServerTransport; import org.apache.thrift.transport.TSocket; import org.apache.thrift.transport.TTransportException; +import org.apache.thrift.transport.layered.TFramedTransport; import java.io.IOException; import java.net.InetSocketAddress; @@ -155,13 +157,27 @@ private void createThreadPoolServer() throws TTransportException { .backlog(Config.thrift_backlog_num); } - TThreadPoolServer.Args serverArgs = new TThreadPoolServer.Args(new ImprovedTServerSocket(socketTransportArgs)) + server = createThreadPoolServer(new ImprovedTServerSocket(socketTransportArgs), processor, false); + } + + /** + * Creates a blocking Thrift server that uses Doris' standard daemon worker pool. + * + *

Callers that interoperate with a threaded-selector endpoint can retain its framed wire + * protocol by setting {@code useFramedTransport} to true while using a blocking server transport. + */ + public static TThreadPoolServer createThreadPoolServer( + TServerTransport serverTransport, TProcessor processor, boolean useFramedTransport) { + TThreadPoolServer.Args serverArgs = new TThreadPoolServer.Args(serverTransport) .protocolFactory(new TBinaryProtocol.Factory()) .processor(processor); + if (useFramedTransport) { + serverArgs.transportFactory(new TFramedTransport.Factory(Config.thrift_max_frame_size)); + } ThreadPoolExecutor threadPoolExecutor = ThreadPoolManager.newDaemonCacheThreadPool( Config.thrift_server_max_worker_threads, "thrift-server-pool", true); serverArgs.executorService(threadPoolExecutor); - server = new TThreadPoolServer(serverArgs); + return new TThreadPoolServer(serverArgs); } public void start() throws IOException { diff --git a/fe/fe-core/src/test/java/org/apache/doris/common/ThriftServerTest.java b/fe/fe-core/src/test/java/org/apache/doris/common/ThriftServerTest.java new file mode 100644 index 00000000000000..544f51980126f6 --- /dev/null +++ b/fe/fe-core/src/test/java/org/apache/doris/common/ThriftServerTest.java @@ -0,0 +1,99 @@ +// 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.doris.common; + +import org.apache.thrift.TProcessor; +import org.apache.thrift.protocol.TBinaryProtocol; +import org.apache.thrift.protocol.TProtocol; +import org.apache.thrift.server.ServerContext; +import org.apache.thrift.server.TServerEventHandler; +import org.apache.thrift.server.TThreadPoolServer; +import org.apache.thrift.transport.TServerSocket; +import org.apache.thrift.transport.TSocket; +import org.apache.thrift.transport.TTransport; +import org.apache.thrift.transport.layered.TFramedTransport; +import org.junit.Assert; +import org.junit.Test; + +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; + +public class ThriftServerTest { + @Test + public void testFramedThreadPoolRoundTrip() throws Exception { + assertThreadPoolRoundTrip(true); + } + + @Test + public void testUnframedThreadPoolRoundTrip() throws Exception { + assertThreadPoolRoundTrip(false); + } + + private void assertThreadPoolRoundTrip(boolean useFramedTransport) throws Exception { + TServerSocket serverTransport = new TServerSocket(0); + int port = serverTransport.getServerSocket().getLocalPort(); + TProcessor processor = (input, output) -> { + int request = input.readI32(); + output.writeI32(request + 1); + output.getTransport().flush(); + }; + TThreadPoolServer server = ThriftServer.createThreadPoolServer( + serverTransport, processor, useFramedTransport); + CountDownLatch started = new CountDownLatch(1); + server.setServerEventHandler(new TServerEventHandler() { + @Override + public void preServe() { + started.countDown(); + } + + @Override + public ServerContext createContext(TProtocol input, TProtocol output) { + return null; + } + + @Override + public void deleteContext(ServerContext serverContext, TProtocol input, TProtocol output) { + } + + @Override + public void processContext( + ServerContext serverContext, TTransport inputTransport, TTransport outputTransport) { + } + }); + + Thread serverThread = new Thread(server::serve, "framed-thrift-server-test"); + serverThread.setDaemon(true); + serverThread.start(); + Assert.assertTrue(started.await(5, TimeUnit.SECONDS)); + + TSocket clientSocket = new TSocket("127.0.0.1", port, 5000); + TTransport clientTransport = useFramedTransport ? new TFramedTransport(clientSocket) : clientSocket; + try { + clientTransport.open(); + TBinaryProtocol protocol = new TBinaryProtocol(clientTransport); + protocol.writeI32(41); + clientTransport.flush(); + Assert.assertEquals(42, protocol.readI32()); + } finally { + clientTransport.close(); + server.stop(); + serverThread.join(5000); + } + Assert.assertFalse(serverThread.isAlive()); + } +}