ref[net]: use zero copy in tunnel proxy server

This commit is contained in:
jaysunxiao committed 2025-06-13 15:53:35 +08:00
1 parent 7a8395b951
commit 21fed98235
11 files changed
+159 -56

No files matched your search

@@ -15,6 +15,7 @@ package com.zfoo.net.core.proxy;
import com.zfoo.net.core.AbstractClient;
import com.zfoo.net.core.HostAndPort;
import com.zfoo.net.core.proxy.handler.TunnelClientIdleHandler;
import com.zfoo.net.core.proxy.handler.TunnelClientRouteHandler;
import com.zfoo.net.core.proxy.handler.TunnelClientCodecHandler;
import com.zfoo.net.handler.idle.ClientIdleHandler;
@@ -39,7 +40,7 @@ public class TunnelClient extends AbstractClient<SocketChannel> {
@Override
protected void initChannel(SocketChannel channel) {
channel.pipeline().addLast(new IdleStateHandler(0, 0, 60));
channel.pipeline().addLast(new ClientIdleHandler());
channel.pipeline().addLast(new TunnelClientIdleHandler());
channel.pipeline().addLast(new TunnelClientCodecHandler());
channel.pipeline().addLast(new TunnelClientRouteHandler());
}
@@ -15,9 +15,11 @@ package com.zfoo.net.core.proxy;
import com.zfoo.net.NetContext;
import com.zfoo.net.packet.EncodedPacketInfo;
import com.zfoo.net.packet.PacketService;
import com.zfoo.net.util.SessionUtils;
import com.zfoo.protocol.buffer.ByteBufUtils;
import io.netty.buffer.ByteBuf;
import io.netty.channel.Channel;
import io.netty.util.ReferenceCountUtil;
/**
@@ -33,6 +35,7 @@ public class TunnelProtocolClient2Server {
public static final byte FLAG_PACKET = 0;
public static final byte FLAG_REGISTER = 10;
public static final byte FLAG_HEARTBEAT = 20;
private long sid;
@@ -85,28 +88,41 @@ public class TunnelProtocolClient2Server {
}
public static void writeRegister(ByteBuf out, TunnelRegister register) {
// out.ensureWritable(4);
// out.writerIndex(PacketService.PACKET_HEAD_LENGTH);
// out.writeByte(FLAG_PACKET);
// ByteBufUtils.writeLong(out, encodedPacketInfo.getSid());
// ByteBufUtils.writeLong(out, encodedPacketInfo.getUid());
//
// var packet = encodedPacketInfo.getPacket();
// var attachment = encodedPacketInfo.getAttachment();
// NetContext.getPacketService().write(out, packet, attachment);
// NetContext.getPacketService().writeHeaderBefore(out);
out.ensureWritable(4);
out.writerIndex(PacketService.PACKET_HEAD_LENGTH);
out.writeByte(FLAG_REGISTER);
ByteBufUtils.writeLong(out, register.sid);
NetContext.getPacketService().writeHeaderBefore(out);
}
// -----------------------------------------------------------------------------------------------------------------
public static class TunnelHeartbeat {
}
public static void writeHeartbeat(ByteBuf out, TunnelHeartbeat heartbeat) {
out.ensureWritable(4);
out.writerIndex(PacketService.PACKET_HEAD_LENGTH);
out.writeByte(FLAG_HEARTBEAT);
NetContext.getPacketService().writeHeaderBefore(out);
}
// -----------------------------------------------------------------------------------------------------------------
public static void read(Channel channel, ByteBuf in) {
var flag = in.readByte();
if (flag == FLAG_REGISTER) {
TunnelServer.tunnels.add(channel);
} else if (flag == FLAG_PACKET) {
if (flag == FLAG_PACKET) {
var sid = ByteBufUtils.readLong(in);
var uid = ByteBufUtils.readLong(in);
var session = NetContext.getSessionManager().getServerSession(sid);
session.getChannel().writeAndFlush(TunnelProtocolServer2Client.valueOf(sid, uid, in));
if (SessionUtils.isActive(session)) {
session.getChannel().writeAndFlush(TunnelProtocolServer2Client.valueOf(sid, uid, in));
} else {
ReferenceCountUtil.release(in);
}
} else if (flag == FLAG_REGISTER) {
TunnelServer.tunnels.add(channel);
ReferenceCountUtil.release(in);
} else if (flag == FLAG_HEARTBEAT) {
ReferenceCountUtil.release(in);
}
}
@@ -16,6 +16,7 @@ import com.zfoo.net.NetContext;
import com.zfoo.net.packet.PacketService;
import com.zfoo.protocol.buffer.ByteBufUtils;
import io.netty.buffer.ByteBuf;
import io.netty.util.ReferenceCountUtil;
/**
@@ -27,13 +28,13 @@ public class TunnelProtocolServer2Client {
private long uid;
private ByteBuf byteBuf;
private ByteBuf retainedByteBuf;
public static TunnelProtocolServer2Client valueOf(long sid, long uid, ByteBuf byteBuf) {
public static TunnelProtocolServer2Client valueOf(long sid, long uid, ByteBuf retainedByteBuf) {
var tunnelProtocol = new TunnelProtocolServer2Client();
tunnelProtocol.sid = sid;
tunnelProtocol.uid = uid;
tunnelProtocol.byteBuf = byteBuf;
tunnelProtocol.retainedByteBuf = retainedByteBuf;
return tunnelProtocol;
}
@@ -66,12 +67,16 @@ public class TunnelProtocolServer2Client {
public void write(ByteBuf out) {
out.ensureWritable(4);
out.writerIndex(PacketService.PACKET_HEAD_LENGTH);
ByteBufUtils.writeLong(out, sid);
ByteBufUtils.writeLong(out, uid);
out.writeBytes(byteBuf);
NetContext.getPacketService().writeHeaderBefore(out);
try {
out.ensureWritable(22);
out.writerIndex(PacketService.PACKET_HEAD_LENGTH);
ByteBufUtils.writeLong(out, sid);
ByteBufUtils.writeLong(out, uid);
out.writeBytes(retainedByteBuf);
NetContext.getPacketService().writeHeaderBefore(out);
} finally {
ReferenceCountUtil.release(retainedByteBuf);
}
}
// -----------------------------------------------------------------------------------------------------------------
@@ -84,7 +89,7 @@ public class TunnelProtocolServer2Client {
return uid;
}
public ByteBuf getByteBuf() {
return byteBuf;
public ByteBuf getRetainedByteBuf() {
return retainedByteBuf;
}
}
@@ -18,12 +18,14 @@ import com.zfoo.net.core.proxy.TunnelProtocolServer2Client;
import com.zfoo.net.core.proxy.TunnelServer;
import com.zfoo.net.packet.PacketService;
import com.zfoo.net.util.SessionUtils;
import com.zfoo.protocol.collection.CollectionUtils;
import com.zfoo.protocol.util.IOUtils;
import com.zfoo.protocol.util.RandomUtils;
import com.zfoo.protocol.util.StringUtils;
import io.netty.buffer.ByteBuf;
import io.netty.channel.ChannelHandlerContext;
import io.netty.handler.codec.ByteToMessageCodec;
import io.netty.util.ReferenceCountUtil;
import java.util.List;
@@ -53,19 +55,35 @@ public class ProxyCodecHandler extends ByteToMessageCodec<TunnelProtocolServer2C
return;
}
var sliceByteBuf = in.readSlice(length);
var session = SessionUtils.getSession(ctx);
if (CollectionUtils.isEmpty(TunnelServer.tunnels)) {
in.readSlice(length);
return;
}
var tunnel = RandomUtils.randomEle(TunnelServer.tunnels);
tunnel.writeAndFlush(TunnelProtocolServer2Client.valueOf(session.getSid(), session.getUid(), sliceByteBuf));
if (!SessionUtils.isActive(tunnel)) {
in.readSlice(length);
return;
}
var retainedByteBuf = in.readRetainedSlice(length);
try {
var session = SessionUtils.getSession(ctx);
tunnel.writeAndFlush(TunnelProtocolServer2Client.valueOf(session.getSid(), session.getUid(), retainedByteBuf));
} catch (Throwable t) {
ReferenceCountUtil.release(retainedByteBuf);
}
}
@Override
protected void encode(ChannelHandlerContext ctx, TunnelProtocolServer2Client tunnelProtocol, ByteBuf out) {
out.ensureWritable(7);
out.writerIndex(PacketService.PACKET_HEAD_LENGTH);
out.writeBytes(tunnelProtocol.getByteBuf());
NetContext.getPacketService().writeHeaderBefore(out);
try {
out.ensureWritable(7);
out.writerIndex(PacketService.PACKET_HEAD_LENGTH);
out.writeBytes(tunnelProtocol.getRetainedByteBuf());
NetContext.getPacketService().writeHeaderBefore(out);
} finally {
ReferenceCountUtil.release(tunnelProtocol.getRetainedByteBuf());
}
}
}
@@ -22,8 +22,6 @@ import com.zfoo.protocol.util.StringUtils;
import io.netty.buffer.ByteBuf;
import io.netty.channel.ChannelHandlerContext;
import io.netty.handler.codec.ByteToMessageCodec;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import java.util.List;
@@ -63,6 +61,10 @@ public class TunnelClientCodecHandler extends ByteToMessageCodec<Object> {
protected void encode(ChannelHandlerContext ctx, Object msg, ByteBuf out) {
if (msg instanceof EncodedPacketInfo) {
TunnelProtocolClient2Server.writePacket(out, (EncodedPacketInfo) msg);
} else if (msg instanceof TunnelProtocolClient2Server.TunnelHeartbeat) {
TunnelProtocolClient2Server.writeHeartbeat(out, (TunnelProtocolClient2Server.TunnelHeartbeat) msg);
} else if (msg instanceof TunnelProtocolClient2Server.TunnelRegister) {
TunnelProtocolClient2Server.writeRegister(out, (TunnelProtocolClient2Server.TunnelRegister) msg);
}
}
@@ -0,0 +1,44 @@
/*
* Copyright (C) 2020 The zfoo 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
*
* 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 com.zfoo.net.core.proxy.handler;
import com.zfoo.net.core.proxy.TunnelProtocolClient2Server;
import com.zfoo.net.util.SessionUtils;
import io.netty.channel.ChannelDuplexHandler;
import io.netty.channel.ChannelHandlerContext;
import io.netty.handler.timeout.IdleState;
import io.netty.handler.timeout.IdleStateEvent;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
/**
* @author jaysunxiao
*/
public class TunnelClientIdleHandler extends ChannelDuplexHandler {
private static final Logger logger = LoggerFactory.getLogger(TunnelClientIdleHandler.class);
private static final TunnelProtocolClient2Server.TunnelHeartbeat heartbeat = new TunnelProtocolClient2Server.TunnelHeartbeat();
@Override
public void userEventTriggered(ChannelHandlerContext ctx, Object evt) {
if (evt instanceof IdleStateEvent) {
IdleStateEvent event = (IdleStateEvent) evt;
if (event.state() == IdleState.ALL_IDLE) {
logger.info("client send heartbeat to [sid:{}]", SessionUtils.getSession(ctx).getSid());
ctx.channel().writeAndFlush(heartbeat);
}
}
}
}
@@ -19,6 +19,7 @@ import com.zfoo.net.core.proxy.TunnelProtocolClient2Server;
import com.zfoo.net.core.proxy.TunnelProtocolServer2Client;
import com.zfoo.net.handler.ClientRouteHandler;
import com.zfoo.net.session.Session;
import com.zfoo.net.util.SessionUtils;
import io.netty.channel.ChannelHandler;
import io.netty.channel.ChannelHandlerContext;
@@ -32,8 +33,8 @@ public class TunnelClientRouteHandler extends ClientRouteHandler {
public void channelActive(ChannelHandlerContext ctx) throws Exception {
super.channelActive(ctx);
TunnelClient.tunnels.add(ctx.channel());
ctx.channel().writeAndFlush(new TunnelProtocolClient2Server.TunnelRegister(1));
var session = SessionUtils.getSession(ctx);
ctx.channel().writeAndFlush(new TunnelProtocolClient2Server.TunnelRegister(session.getSid()));
}
@Override
@@ -21,8 +21,7 @@ import com.zfoo.protocol.util.StringUtils;
import io.netty.buffer.ByteBuf;
import io.netty.channel.ChannelHandlerContext;
import io.netty.handler.codec.ByteToMessageCodec;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import io.netty.util.ReferenceCountUtil;
import java.util.List;
@@ -31,8 +30,6 @@ import java.util.List;
*/
public class TunnelServerCodecHandler extends ByteToMessageCodec<TunnelProtocolServer2Client> {
private static final Logger logger = LoggerFactory.getLogger(TunnelServerCodecHandler.class);
@Override
protected void decode(ChannelHandlerContext ctx, ByteBuf in, List<Object> out) {
// 不够读一个int
@@ -53,8 +50,12 @@ public class TunnelServerCodecHandler extends ByteToMessageCodec<TunnelProtocolS
return;
}
var sliceByteBuf = in.readSlice(length);
TunnelProtocolClient2Server.read(ctx.channel(), sliceByteBuf);
var retainedByteBuf = in.readRetainedSlice(length);
try {
TunnelProtocolClient2Server.read(ctx.channel(), retainedByteBuf);
} catch (Throwable t) {
ReferenceCountUtil.release(retainedByteBuf);
}
}
@Override
@@ -16,8 +16,8 @@ package com.zfoo.net.core.proxy.client;
import com.zfoo.net.NetContext;
import com.zfoo.net.core.HostAndPort;
import com.zfoo.net.core.tcp.TcpClient;
import com.zfoo.net.packet.tcp.TcpHelloRequest;
import com.zfoo.net.packet.tcp.TcpHelloResponse;
import com.zfoo.net.packet.proxy.ProxyHelloRequest;
import com.zfoo.net.packet.proxy.ProxyHelloResponse;
import com.zfoo.protocol.util.JsonUtils;
import com.zfoo.protocol.util.ThreadUtils;
import org.junit.Ignore;
@@ -26,6 +26,8 @@ import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.context.support.ClassPathXmlApplicationContext;
import java.util.function.Consumer;
/**
* @author jaysunxiao
*/
@@ -41,9 +43,22 @@ public class ProxyClientTest {
var client = new TcpClient(HostAndPort.valueOf("127.0.0.1:9000"));
var session = client.start();
var request = TcpHelloRequest.valueOf("Hello, this is the tcp client!");
var response = NetContext.getRouter().syncAsk(session, request, TcpHelloResponse.class, null).packet();
logger.info("sync client receive [packet:{}] from server", JsonUtils.object2String(response));
var request = ProxyHelloRequest.valueOf("Hello, this is the tcp client!");
for (int i = 0; i < 1000; i++) {
NetContext.getRouter().send(session, request);
var response = NetContext.getRouter().syncAsk(session, request, ProxyHelloResponse.class, null).packet();
logger.info("sync client receive [packet:{}] from server", JsonUtils.object2String(response));
NetContext.getRouter().asyncAsk(session, request, ProxyHelloResponse.class, null)
.whenComplete(new Consumer<ProxyHelloResponse>() {
@Override
public void accept(ProxyHelloResponse jsonHelloResponse) {
logger.info("async client receive [packet:{}] from server", JsonUtils.object2String(jsonHelloResponse));
}
});
}
ThreadUtils.sleep(Long.MAX_VALUE);
}
@@ -15,7 +15,7 @@ package com.zfoo.net.core.proxy.server;
import com.zfoo.net.NetContext;
import com.zfoo.net.anno.PacketReceiver;
import com.zfoo.net.packet.proxy.ProxyHelloRequest;
import com.zfoo.net.packet.tcp.TcpHelloRequest;
import com.zfoo.net.packet.proxy.ProxyHelloResponse;
import com.zfoo.net.packet.tcp.TcpHelloResponse;
import com.zfoo.net.session.Session;
import com.zfoo.protocol.util.JsonUtils;
@@ -27,15 +27,15 @@ import org.springframework.stereotype.Component;
* @author jaysunxiao
*/
@Component
public class ReverseProxyServerController {
public class ProxyServerController {
private static final Logger logger = LoggerFactory.getLogger(ReverseProxyServerController.class);
private static final Logger logger = LoggerFactory.getLogger(ProxyServerController.class);
@PacketReceiver
public void atProxyHelloRequest(Session session, ProxyHelloRequest request) {
logger.info("receive [packet:{}] from client", JsonUtils.object2String(request));
var response = new TcpHelloResponse();
var response = new ProxyHelloResponse();
response.setMessage("Hello, this is the proxy server! -> " + request.getMessage());
NetContext.getRouter().send(session, response);
@@ -14,8 +14,8 @@
package com.zfoo.net.core.proxy.server;
import com.zfoo.net.core.HostAndPort;
import com.zfoo.net.core.proxy.ProxyTcpServer;
import com.zfoo.net.core.proxy.TunnelServer;
import com.zfoo.net.core.tcp.TcpServer;
import com.zfoo.protocol.util.ThreadUtils;
import org.junit.Ignore;
import org.junit.Test;
@@ -34,8 +34,8 @@ public class ProxyServerTest {
public void startServer() {
var context = new ClassPathXmlApplicationContext("config.xml");
var server = new TcpServer(HostAndPort.valueOf("0.0.0.0:9000"));
server.start();
var proxyServer = new ProxyTcpServer(HostAndPort.valueOf("0.0.0.0:9000"));
proxyServer.start();
var tunnelServer = new TunnelServer(HostAndPort.valueOf("0.0.0.0:9001"));
tunnelServer.start();