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
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ public abstract class ConnListenerBuilder<C extends ConnListenerBuilder<C>> {
private final MQTTBrokerBuilder serverBuilder;
protected String host;
protected int port;
protected boolean enableProxyProtocol = true;

ConnListenerBuilder(MQTTBrokerBuilder builder) {
serverBuilder = builder;
Expand Down Expand Up @@ -63,6 +64,11 @@ public C port(int port) {
return thisT();
}

public C enableProxyProtocol(boolean enableProxyProtocol) {
this.enableProxyProtocol = enableProxyProtocol;
return thisT();
}

public <T> C option(ChannelOption<T> option, T value) {
Preconditions.checkNotNull(option, "option");
if (value == null) {
Expand Down Expand Up @@ -120,6 +126,7 @@ public static final class TLSConnListenerBuilder extends SecuredConnListenerBuil

public static final class WSConnListenerBuilder extends ConnListenerBuilder<WSConnListenerBuilder> {
private String path = "mqtt";
private boolean enableClientAddressHeader = true;

WSConnListenerBuilder(MQTTBrokerBuilder builder) {
super(builder);
Expand All @@ -133,10 +140,20 @@ public WSConnListenerBuilder path(String path) {
this.path = path;
return this;
}

public WSConnListenerBuilder enableClientAddressHeader(boolean enableClientAddressHeader) {
this.enableClientAddressHeader = enableClientAddressHeader;
return this;
}

public boolean clientAddressHeaderEnabled() {
return enableClientAddressHeader;
}
}

public static final class WSSConnListenerBuilder extends SecuredConnListenerBuilder<WSSConnListenerBuilder> {
private String path;
private boolean enableClientAddressHeader = true;

WSSConnListenerBuilder(MQTTBrokerBuilder builder) {
super(builder);
Expand All @@ -150,5 +167,14 @@ public WSSConnListenerBuilder path(String path) {
this.path = path;
return this;
}

public WSSConnListenerBuilder enableClientAddressHeader(boolean enableClientAddressHeader) {
this.enableClientAddressHeader = enableClientAddressHeader;
return this;
}

public boolean clientAddressHeaderEnabled() {
return enableClientAddressHeader;
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -170,7 +170,7 @@ public final void close() {
}

private ChannelFuture bindTCPChannel(ConnListenerBuilder.TCPConnListenerBuilder connBuilder) {
return buildChannel(connBuilder, new MQTTChannelInitializer() {
return buildChannel(connBuilder, new MQTTChannelInitializer(connBuilder.enableProxyProtocol) {
@Override
protected void initChannel(SocketChannel ch) {
super.initChannel(ch);
Expand All @@ -194,7 +194,7 @@ protected void initChannel(SocketChannel ch) {
}

private ChannelFuture bindTLSChannel(ConnListenerBuilder.TLSConnListenerBuilder connBuilder) {
return buildChannel(connBuilder, new MQTTChannelInitializer() {
return buildChannel(connBuilder, new MQTTChannelInitializer(connBuilder.enableProxyProtocol) {
@Override
protected void initChannel(SocketChannel ch) {
super.initChannel(ch);
Expand All @@ -219,7 +219,7 @@ protected void initChannel(SocketChannel ch) {
}

private ChannelFuture bindWSChannel(ConnListenerBuilder.WSConnListenerBuilder connBuilder) {
return buildChannel(connBuilder, new MQTTChannelInitializer() {
return buildChannel(connBuilder, new MQTTChannelInitializer(connBuilder.enableProxyProtocol) {
@Override
protected void initChannel(SocketChannel ch) {
super.initChannel(ch);
Expand All @@ -229,7 +229,9 @@ protected void initChannel(SocketChannel ch) {
new ChannelTrafficShapingHandler(builder.writeLimit, builder.readLimit));
p.addLast("httpEncoder", new HttpResponseEncoder());
p.addLast("httpDecoder", new HttpRequestDecoder());
p.addLast("remoteAddr", new ClientAddrHandler());
if (connBuilder.clientAddressHeaderEnabled()) {
p.addLast("remoteAddr", new ClientAddrHandler());
}
p.addLast("aggregator", new HttpObjectAggregator(65536));
p.addLast("webSocketOnly", new WebSocketOnlyHandler(connBuilder.path()));
p.addLast("webSocketHandler", new WebSocketServerProtocolHandler(connBuilder.path(),
Expand All @@ -242,7 +244,7 @@ protected void initChannel(SocketChannel ch) {
}

private ChannelFuture bindWSSChannel(ConnListenerBuilder.WSSConnListenerBuilder connBuilder) {
return buildChannel(connBuilder, new MQTTChannelInitializer() {
return buildChannel(connBuilder, new MQTTChannelInitializer(connBuilder.enableProxyProtocol) {
@Override
protected void initChannel(SocketChannel ch) {
super.initChannel(ch);
Expand All @@ -253,7 +255,9 @@ protected void initChannel(SocketChannel ch) {
new ChannelTrafficShapingHandler(builder.writeLimit, builder.readLimit));
p.addLast("httpEncoder", new HttpResponseEncoder());
p.addLast("httpDecoder", new HttpRequestDecoder());
p.addLast(ClientAddrHandler.class.getName(), new ClientAddrHandler());
if (connBuilder.clientAddressHeaderEnabled()) {
p.addLast(ClientAddrHandler.class.getName(), new ClientAddrHandler());
}
p.addLast("aggregator", new HttpObjectAggregator(65536));
p.addLast("webSocketOnly", new WebSocketOnlyHandler(connBuilder.path()));
p.addLast("webSocketHandler", new WebSocketServerProtocolHandler(connBuilder.path(),
Expand All @@ -279,14 +283,22 @@ private <T extends ConnListenerBuilder<T>> ChannelFuture buildChannel(T builder,
}

private abstract static class MQTTChannelInitializer extends ChannelInitializer<SocketChannel> {
private final boolean enableProxyProtocol;

private MQTTChannelInitializer(boolean enableProxyProtocol) {
this.enableProxyProtocol = enableProxyProtocol;
}

@Override
protected void initChannel(SocketChannel ch) {
ChannelPipeline pipeline = ch.pipeline();
// handler for proxy protocol v1 and v2
pipeline
.addLast(ProxyProtocolDetector.class.getName(), new ProxyProtocolDetector())
.addLast(HAProxyMessageDecoder.class.getName(), new HAProxyMessageDecoder())
.addLast(ProxyProtocolHandler.class.getName(), new ProxyProtocolHandler());
if (enableProxyProtocol) {
ChannelPipeline pipeline = ch.pipeline();
// handler for proxy protocol v1 and v2
pipeline
.addLast(ProxyProtocolDetector.class.getName(), new ProxyProtocolDetector())
.addLast(HAProxyMessageDecoder.class.getName(), new HAProxyMessageDecoder())
.addLast(ProxyProtocolHandler.class.getName(), new ProxyProtocolHandler());
}
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -28,4 +28,5 @@ public class TCPListenerConfig {
private boolean enable = true;
private String host = "0.0.0.0";
private int port = 1883;
private boolean enableProxyProtocol = true;
}
Original file line number Diff line number Diff line change
Expand Up @@ -30,4 +30,5 @@ public class TLSListenerConfig {
private String host;
private int port = 1884;
private ServerSSLContextConfig sslConfig;
private boolean enableProxyProtocol = true;
}
Original file line number Diff line number Diff line change
Expand Up @@ -29,4 +29,6 @@ public class WSListenerConfig {
private String host;
private int port = 8080;
private String wsPath = "/mqtt";
private boolean enableProxyProtocol = true;
private boolean enableClientAddressHeader = true;
}
Original file line number Diff line number Diff line change
Expand Up @@ -31,4 +31,6 @@ public class WSSListenerConfig {
private int port = 8443;
private String wsPath = "/mqtt";
private ServerSSLContextConfig sslConfig;
private boolean enableProxyProtocol = true;
private boolean enableClientAddressHeader = true;
}
Original file line number Diff line number Diff line change
Expand Up @@ -89,12 +89,14 @@ public Optional<IMQTTBroker> get() {
brokerBuilder.buildTcpConnListener()
.host(serverConfig.getTcpListener().getHost())
.port(serverConfig.getTcpListener().getPort())
.enableProxyProtocol(serverConfig.getTcpListener().isEnableProxyProtocol())
.buildListener();
}
if (serverConfig.getTlsListener().isEnable()) {
brokerBuilder.buildTLSConnListener()
.host(serverConfig.getTlsListener().getHost())
.port(serverConfig.getTlsListener().getPort())
.enableProxyProtocol(serverConfig.getTlsListener().isEnableProxyProtocol())
.sslContext(buildServerSslContext(serverConfig.getTlsListener().getSslConfig()))
.buildListener();
}
Expand All @@ -103,13 +105,17 @@ public Optional<IMQTTBroker> get() {
.host(serverConfig.getWsListener().getHost())
.port(serverConfig.getWsListener().getPort())
.path(serverConfig.getWsListener().getWsPath())
.enableProxyProtocol(serverConfig.getWsListener().isEnableProxyProtocol())
.enableClientAddressHeader(serverConfig.getWsListener().isEnableClientAddressHeader())
.buildListener();
}
if (serverConfig.getWssListener().isEnable()) {
brokerBuilder.buildWSSConnListener()
.host(serverConfig.getWssListener().getHost())
.port(serverConfig.getWssListener().getPort())
.path(serverConfig.getWssListener().getWsPath())
.enableProxyProtocol(serverConfig.getWssListener().isEnableProxyProtocol())
.enableClientAddressHeader(serverConfig.getWssListener().isEnableClientAddressHeader())
.sslContext(buildServerSslContext(serverConfig.getWssListener().getSslConfig()))
.buildListener();
}
Expand Down
Loading