This is an automated email from the ASF dual-hosted git repository.
Aias00 pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/shenyu.git
The following commit(s) were added to refs/heads/master by this push:
new b0b7fab03b fix: enforce per-channel MQTT connection state guard (#6983)
b0b7fab03b is described below
commit b0b7fab03b198c561c6f9f363755ff8eae9b1913
Author: wy471x <[email protected]>
AuthorDate: Fri Sep 4 11:20:53 2026 +0800
fix: enforce per-channel MQTT connection state guard (#6983)
* fix: enforce mqtt connection state guard per channel to reject
pre-connect operations and duplicate connect
Co-Authored-By: Claude <[email protected]>
* fix: clean up channel repository entry when mqtt channel becomes inactive
A channel registered by a successful CONNECT stayed in ChannelRepository
when it was closed outside the explicit DISCONNECT flow, e.g. after a
duplicate CONNECT rejected by the connection state guard or an abrupt
client disconnect. Clean up centrally in channelInactive, and make the
registration synchronous so an in-flight put can not re-add a closed
channel after the removal.
---------
Co-authored-by: Claude <[email protected]>
Co-authored-by: aias00 <[email protected]>
---
.../org/apache/shenyu/protocol/mqtt/Connect.java | 8 ++-
.../apache/shenyu/protocol/mqtt/MessageType.java | 14 ++--
.../shenyu/protocol/mqtt/MqttTransportHandler.java | 8 +++
.../org/apache/shenyu/protocol/mqtt/PingReq.java | 6 ++
.../org/apache/shenyu/protocol/mqtt/Publish.java | 4 +-
.../org/apache/shenyu/protocol/mqtt/Subscribe.java | 2 +-
.../apache/shenyu/protocol/mqtt/Unsubscribe.java | 4 +-
.../mqtt/repositories/ChannelRepository.java | 3 +-
.../apache/shenyu/protocol/mqtt/ConnectTest.java | 54 +++++++++------
...nectTest.java => MqttTransportHandlerTest.java} | 81 ++++++++--------------
.../apache/shenyu/protocol/mqtt/PingReqTest.java | 56 +++++++++++++++
.../apache/shenyu/protocol/mqtt/PublishTest.java | 78 +++++++++++++++++++--
.../shenyu/protocol/mqtt/UnsubscribeTest.java | 58 ++++++++++++++++
13 files changed, 284 insertions(+), 92 deletions(-)
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Connect.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Connect.java
index bca27d001d..b23b07a661 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Connect.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Connect.java
@@ -44,6 +44,12 @@ public class Connect extends MessageType {
@Override
public void connect(final ChannelHandlerContext ctx, final
MqttConnectMessage msg) {
+ if (isConnected(ctx.channel())) {
+ LOG.info("MQTT client has already sent a CONNECT packet, closing
connection.");
+ ctx.close().addListener(CLOSE_ON_FAILURE);
+ return;
+ }
+
String clientId = msg.payload().clientIdentifier();
if (StringUtils.isEmpty(clientId)) {
LOG.info("MQTT clientId can not be empty.");
@@ -73,7 +79,7 @@ public class Connect extends MessageType {
.sessionPresent(true)
.build();
ctx.writeAndFlush(ackMessage);
- setConnected(true);
+ setConnected(ctx.channel(), true);
}
private void close(final ChannelHandlerContext ctx, final
MqttConnectReturnCode returnCode) {
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MessageType.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MessageType.java
index 9d766f31bc..f6bdecac39 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MessageType.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MessageType.java
@@ -17,33 +17,37 @@
package org.apache.shenyu.protocol.mqtt;
+import io.netty.channel.Channel;
import io.netty.channel.ChannelHandlerContext;
import io.netty.handler.codec.mqtt.MqttConnectMessage;
import io.netty.handler.codec.mqtt.MqttPublishMessage;
import io.netty.handler.codec.mqtt.MqttSubscribeMessage;
import io.netty.handler.codec.mqtt.MqttUnsubscribeMessage;
+import io.netty.util.AttributeKey;
/**
* Command messages.
*/
public class MessageType implements AbstractMessageType {
- private volatile boolean connected;
+ private static final AttributeKey<Boolean> CONNECTED =
AttributeKey.valueOf("connected");
/**
* isConnected.
+ * @param channel channel
* @return connected
*/
- boolean isConnected() {
- return connected;
+ protected boolean isConnected(final Channel channel) {
+ return Boolean.TRUE.equals(channel.attr(CONNECTED).get());
}
/**
* set connected.
+ * @param channel channel
* @param connected connected
*/
- void setConnected(final boolean connected) {
- this.connected = connected;
+ protected void setConnected(final Channel channel, final boolean
connected) {
+ channel.attr(CONNECTED).set(connected);
}
@Override
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java
index 49ddc8a058..a6fa975297 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java
@@ -22,6 +22,8 @@ import io.netty.channel.ChannelInboundHandlerAdapter;
import io.netty.handler.codec.mqtt.MqttMessage;
import io.netty.util.concurrent.Future;
import io.netty.util.concurrent.GenericFutureListener;
+import org.apache.shenyu.common.utils.Singleton;
+import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
/**
* mqtt transport handler.
@@ -38,6 +40,12 @@ public class MqttTransportHandler extends
ChannelInboundHandlerAdapter implement
}
}
+ @Override
+ public void channelInactive(final ChannelHandlerContext ctx) throws
Exception {
+ Singleton.INST.get(ChannelRepository.class).remove(ctx.channel());
+ ctx.fireChannelInactive();
+ }
+
@Override
public void operationComplete(final Future<? super Void> future) throws
Exception {
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PingReq.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PingReq.java
index 716cf54915..a95db72cb4 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PingReq.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PingReq.java
@@ -19,6 +19,8 @@ package org.apache.shenyu.protocol.mqtt;
import io.netty.channel.ChannelHandlerContext;
+import static io.netty.channel.ChannelFutureListener.FIRE_EXCEPTION_ON_FAILURE;
+
/**
* Client sends pingreq to the server.
*/
@@ -26,6 +28,10 @@ public class PingReq extends MessageType {
@Override
public void pingReq(final ChannelHandlerContext ctx) {
+ if (!isConnected(ctx.channel())) {
+ ctx.channel().close().addListener(FIRE_EXCEPTION_ON_FAILURE);
+ return;
+ }
new PingResp().pingResp(ctx);
}
}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java
index bd7f48ef96..c469b3cf22 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java
@@ -35,6 +35,7 @@ import
org.apache.shenyu.protocol.mqtt.repositories.TopicRepository;
import java.util.List;
import java.util.concurrent.CompletableFuture;
+import static io.netty.channel.ChannelFutureListener.FIRE_EXCEPTION_ON_FAILURE;
import static io.netty.handler.codec.mqtt.MqttMessageType.PUBACK;
/**
@@ -44,7 +45,8 @@ public class Publish extends MessageType {
@Override
public void publish(final ChannelHandlerContext ctx, final
MqttPublishMessage msg) {
- if (isConnected()) {
+ if (!isConnected(ctx.channel())) {
+ ctx.channel().close().addListener(FIRE_EXCEPTION_ON_FAILURE);
return;
}
String topic = msg.variableHeader().topicName();
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Subscribe.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Subscribe.java
index c71a218ca5..1a9975cfd1 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Subscribe.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Subscribe.java
@@ -53,7 +53,7 @@ public class Subscribe extends MessageType {
public void subscribe(final ChannelHandlerContext ctx, final
MqttSubscribeMessage msg) {
Channel channel = ctx.channel();
- if (isConnected()) {
+ if (!isConnected(channel)) {
channel.close().addListener(FIRE_EXCEPTION_ON_FAILURE);
return;
}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Unsubscribe.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Unsubscribe.java
index 376015a6bf..154b61af72 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Unsubscribe.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Unsubscribe.java
@@ -29,6 +29,7 @@ import
org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository;
import java.util.List;
+import static io.netty.channel.ChannelFutureListener.FIRE_EXCEPTION_ON_FAILURE;
import static io.netty.handler.codec.mqtt.MqttMessageIdVariableHeader.from;
/**
@@ -38,7 +39,8 @@ public class Unsubscribe extends MessageType {
@Override
public void unsubscribe(final ChannelHandlerContext ctx, final
MqttUnsubscribeMessage msg) {
- if (isConnected()) {
+ if (!isConnected(ctx.channel())) {
+ ctx.channel().close().addListener(FIRE_EXCEPTION_ON_FAILURE);
return;
}
List<String> topics = msg.payload().topics();
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/ChannelRepository.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/ChannelRepository.java
index decd4c2e06..fe81cd1379 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/ChannelRepository.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/ChannelRepository.java
@@ -20,7 +20,6 @@ package org.apache.shenyu.protocol.mqtt.repositories;
import io.netty.channel.Channel;
import java.util.Map;
-import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ConcurrentHashMap;
/**
@@ -32,7 +31,7 @@ public class ChannelRepository implements
BaseRepository<Channel, String> {
@Override
public void add(final Channel channel, final String clientId) {
- CompletableFuture.runAsync(() -> CHANNEL_FACTORY.put(channel,
clientId));
+ CHANNEL_FACTORY.put(channel, clientId);
}
@Override
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
index 0be443f7c9..2b8acc8e33 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
@@ -17,9 +17,9 @@
package org.apache.shenyu.protocol.mqtt;
-import io.netty.channel.Channel;
-import io.netty.channel.ChannelFuture;
import io.netty.channel.ChannelHandlerContext;
+import io.netty.channel.ChannelInboundHandlerAdapter;
+import io.netty.channel.embedded.EmbeddedChannel;
import io.netty.handler.codec.mqtt.MqttConnAckMessage;
import io.netty.handler.codec.mqtt.MqttConnectMessage;
import io.netty.handler.codec.mqtt.MqttConnectPayload;
@@ -33,7 +33,6 @@ import
org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
import org.junit.jupiter.api.AfterAll;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
-import org.mockito.ArgumentCaptor;
import java.nio.charset.StandardCharsets;
import java.time.Duration;
@@ -42,12 +41,10 @@ import static
io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_ACCEP
import static
io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION;
import static org.awaitility.Awaitility.await;
import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertNotNull;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.junit.jupiter.api.Assertions.assertTrue;
-import static org.mockito.Mockito.mock;
-import static org.mockito.Mockito.times;
-import static org.mockito.Mockito.verify;
-import static org.mockito.Mockito.when;
/**
* Test cases for {@link Connect}.
@@ -93,32 +90,45 @@ public final class ConnectTest {
@Test
public void unsupportedProtocolVersionIsRejected() {
- ChannelHandlerContext ctx = mock(ChannelHandlerContext.class);
- Channel channel = mock(Channel.class);
- when(ctx.channel()).thenReturn(channel);
- when(ctx.close()).thenReturn(mock(ChannelFuture.class));
+ EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ ChannelHandlerContext ctx = channel.pipeline().lastContext();
new Connect().connect(ctx, connectMessage("MQTT", 6));
- ArgumentCaptor<MqttConnAckMessage> captor =
ArgumentCaptor.forClass(MqttConnAckMessage.class);
- verify(ctx, times(1)).writeAndFlush(captor.capture());
+ MqttConnAckMessage ackMessage = channel.readOutbound();
+ assertNotNull(ackMessage);
assertEquals(CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION,
- captor.getValue().variableHeader().connectReturnCode());
- verify(ctx).close();
+ ackMessage.variableHeader().connectReturnCode());
+ channel.runPendingTasks();
+ assertFalse(channel.isActive());
assertNull(channelRepository.get(channel));
}
+ @Test
+ public void duplicateConnectIsRejected() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+ new Connect().connect(ctx,
connectMessage(MqttVersion.MQTT_3_1_1.protocolName(),
MqttVersion.MQTT_3_1_1.protocolLevel()));
+ assertNotNull(channel.readOutbound());
+
+ new Connect().connect(ctx,
connectMessage(MqttVersion.MQTT_3_1_1.protocolName(),
MqttVersion.MQTT_3_1_1.protocolLevel()));
+
+ channel.runPendingTasks();
+ assertFalse(channel.isActive());
+ assertNull(channel.readOutbound());
+ }
+
private void connectIsAccepted(final MqttVersion version) {
- ChannelHandlerContext ctx = mock(ChannelHandlerContext.class);
- Channel channel = mock(Channel.class);
- when(ctx.channel()).thenReturn(channel);
+ EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ ChannelHandlerContext ctx = channel.pipeline().lastContext();
new Connect().connect(ctx, connectMessage(version.protocolName(),
version.protocolLevel()));
- ArgumentCaptor<MqttConnAckMessage> captor =
ArgumentCaptor.forClass(MqttConnAckMessage.class);
- verify(ctx).writeAndFlush(captor.capture());
- assertEquals(CONNECTION_ACCEPTED,
captor.getValue().variableHeader().connectReturnCode());
- assertTrue(captor.getValue().variableHeader().isSessionPresent());
+ MqttConnAckMessage ackMessage = channel.readOutbound();
+ assertNotNull(ackMessage);
+ assertEquals(CONNECTION_ACCEPTED,
ackMessage.variableHeader().connectReturnCode());
+ assertTrue(ackMessage.variableHeader().isSessionPresent());
await().atMost(Duration.ofSeconds(5))
.until(() -> CLIENT_ID.equals(channelRepository.get(channel)));
}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java
similarity index 51%
copy from
shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
copy to
shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java
index 0be443f7c9..2f267af7e1 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java
@@ -17,10 +17,7 @@
package org.apache.shenyu.protocol.mqtt;
-import io.netty.channel.Channel;
-import io.netty.channel.ChannelFuture;
-import io.netty.channel.ChannelHandlerContext;
-import io.netty.handler.codec.mqtt.MqttConnAckMessage;
+import io.netty.channel.embedded.EmbeddedChannel;
import io.netty.handler.codec.mqtt.MqttConnectMessage;
import io.netty.handler.codec.mqtt.MqttConnectPayload;
import io.netty.handler.codec.mqtt.MqttConnectVariableHeader;
@@ -33,26 +30,17 @@ import
org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
import org.junit.jupiter.api.AfterAll;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
-import org.mockito.ArgumentCaptor;
import java.nio.charset.StandardCharsets;
-import java.time.Duration;
-import static
io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_ACCEPTED;
-import static
io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION;
-import static org.awaitility.Awaitility.await;
import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNull;
-import static org.junit.jupiter.api.Assertions.assertTrue;
-import static org.mockito.Mockito.mock;
-import static org.mockito.Mockito.times;
-import static org.mockito.Mockito.verify;
-import static org.mockito.Mockito.when;
/**
- * Test cases for {@link Connect}.
+ * Test cases for {@link MqttTransportHandler}.
*/
-public final class ConnectTest {
+public final class MqttTransportHandlerTest {
private static final String CLIENT_ID = "test-client";
@@ -77,58 +65,43 @@ public final class ConnectTest {
}
@Test
- public void mqtt31ConnectIsAccepted() {
- connectIsAccepted(MqttVersion.MQTT_3_1);
- }
+ public void duplicateConnectCleansUpChannelRepository() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
MqttTransportHandler());
- @Test
- public void mqtt311ConnectIsAccepted() {
- connectIsAccepted(MqttVersion.MQTT_3_1_1);
- }
+ channel.writeInbound(connectMessage());
+ assertEquals(CLIENT_ID, channelRepository.get(channel));
- @Test
- public void mqtt5ConnectIsAccepted() {
- connectIsAccepted(MqttVersion.MQTT_5);
- }
+ channel.writeInbound(connectMessage());
+ channel.runPendingTasks();
- @Test
- public void unsupportedProtocolVersionIsRejected() {
- ChannelHandlerContext ctx = mock(ChannelHandlerContext.class);
- Channel channel = mock(Channel.class);
- when(ctx.channel()).thenReturn(channel);
- when(ctx.close()).thenReturn(mock(ChannelFuture.class));
-
- new Connect().connect(ctx, connectMessage("MQTT", 6));
-
- ArgumentCaptor<MqttConnAckMessage> captor =
ArgumentCaptor.forClass(MqttConnAckMessage.class);
- verify(ctx, times(1)).writeAndFlush(captor.capture());
- assertEquals(CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION,
- captor.getValue().variableHeader().connectReturnCode());
- verify(ctx).close();
+ assertFalse(channel.isActive());
assertNull(channelRepository.get(channel));
+ channel.finishAndReleaseAll();
}
- private void connectIsAccepted(final MqttVersion version) {
- ChannelHandlerContext ctx = mock(ChannelHandlerContext.class);
- Channel channel = mock(Channel.class);
- when(ctx.channel()).thenReturn(channel);
+ @Test
+ public void abruptChannelCloseCleansUpChannelRepository() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
MqttTransportHandler());
+
+ channel.writeInbound(connectMessage());
+ assertEquals(CLIENT_ID, channelRepository.get(channel));
- new Connect().connect(ctx, connectMessage(version.protocolName(),
version.protocolLevel()));
+ channel.close();
+ channel.runPendingTasks();
- ArgumentCaptor<MqttConnAckMessage> captor =
ArgumentCaptor.forClass(MqttConnAckMessage.class);
- verify(ctx).writeAndFlush(captor.capture());
- assertEquals(CONNECTION_ACCEPTED,
captor.getValue().variableHeader().connectReturnCode());
- assertTrue(captor.getValue().variableHeader().isSessionPresent());
- await().atMost(Duration.ofSeconds(5))
- .until(() -> CLIENT_ID.equals(channelRepository.get(channel)));
+ assertFalse(channel.isActive());
+ assertNull(channelRepository.get(channel));
+ channel.finishAndReleaseAll();
}
- private MqttConnectMessage connectMessage(final String protocolName, final
int protocolLevel) {
+ private MqttConnectMessage connectMessage() {
MqttFixedHeader fixedHeader = new
MqttFixedHeader(MqttMessageType.CONNECT, false, MqttQoS.AT_MOST_ONCE, false, 0);
- MqttConnectVariableHeader variableHeader = new
MqttConnectVariableHeader(protocolName, protocolLevel,
+ MqttConnectVariableHeader variableHeader = new
MqttConnectVariableHeader(
+ MqttVersion.MQTT_3_1_1.protocolName(),
MqttVersion.MQTT_3_1_1.protocolLevel(),
true, true, false, 0, false, false, 60);
MqttConnectPayload payload = new MqttConnectPayload(CLIENT_ID, null,
null,
USER_NAME, PASSWORD.getBytes(StandardCharsets.UTF_8));
return new MqttConnectMessage(fixedHeader, variableHeader, payload);
}
+
}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PingReqTest.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PingReqTest.java
new file mode 100644
index 0000000000..9b26ecad9f
--- /dev/null
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PingReqTest.java
@@ -0,0 +1,56 @@
+/*
+ * 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.shenyu.protocol.mqtt;
+
+import io.netty.channel.ChannelHandlerContext;
+import io.netty.channel.ChannelInboundHandlerAdapter;
+import io.netty.channel.embedded.EmbeddedChannel;
+import org.junit.jupiter.api.Test;
+
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertNotNull;
+import static org.junit.jupiter.api.Assertions.assertNull;
+
+/**
+ * Test cases for {@link PingReq}.
+ */
+public final class PingReqTest {
+
+ @Test
+ public void pingReqBeforeConnectClosesChannel() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+ new PingReq().pingReq(ctx);
+
+ channel.runPendingTasks();
+ assertFalse(channel.isActive());
+ assertNull(channel.readOutbound());
+ }
+
+ @Test
+ public void pingReqAfterConnectSendsPingResp() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ new MessageType().setConnected(channel, true);
+ ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+ new PingReq().pingReq(ctx);
+
+ assertNotNull(channel.readOutbound());
+ }
+}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishTest.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishTest.java
index 6d25ff0352..9b5d56ef01 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishTest.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishTest.java
@@ -19,23 +19,32 @@ package org.apache.shenyu.protocol.mqtt;
import io.netty.buffer.Unpooled;
import io.netty.channel.ChannelHandlerContext;
+import io.netty.channel.ChannelInboundHandlerAdapter;
+import io.netty.channel.embedded.EmbeddedChannel;
+import io.netty.handler.codec.mqtt.MqttConnectMessage;
+import io.netty.handler.codec.mqtt.MqttConnectPayload;
+import io.netty.handler.codec.mqtt.MqttConnectVariableHeader;
import io.netty.handler.codec.mqtt.MqttFixedHeader;
import io.netty.handler.codec.mqtt.MqttMessageType;
import io.netty.handler.codec.mqtt.MqttPublishMessage;
import io.netty.handler.codec.mqtt.MqttPublishVariableHeader;
import io.netty.handler.codec.mqtt.MqttQoS;
+import io.netty.handler.codec.mqtt.MqttVersion;
import io.netty.util.CharsetUtil;
import org.apache.shenyu.common.utils.Singleton;
+import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
import org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository;
import org.apache.shenyu.protocol.mqtt.repositories.TopicRepository;
+import org.junit.jupiter.api.AfterAll;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
+import java.nio.charset.StandardCharsets;
import java.time.Duration;
import static org.awaitility.Awaitility.await;
+import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNull;
-import static org.mockito.Mockito.mock;
/**
* Test cases for {@link Publish}.
@@ -48,6 +57,16 @@ public final class PublishTest {
private static final String CLEARED_TOPIC = "test/cleared";
+ private static final String UNCONNECTED_TOPIC = "test/unconnected";
+
+ private static final String END_TO_END_TOPIC = "test/end-to-end";
+
+ private static final String CLIENT_ID = "test-client";
+
+ private static final String USER_NAME = "test-user";
+
+ private static final String PASSWORD = "test-password";
+
private static TopicRepository topicRepository;
@BeforeAll
@@ -55,31 +74,80 @@ public final class PublishTest {
topicRepository = new TopicRepository();
Singleton.INST.single(TopicRepository.class, topicRepository);
Singleton.INST.single(SubscribeRepository.class, new
SubscribeRepository());
+ Singleton.INST.single(ChannelRepository.class, new
ChannelRepository());
+ new MqttContext().setUserName(USER_NAME);
+ new MqttContext().setPassword(PASSWORD);
+ }
+
+ @AfterAll
+ static void tearDown() {
+ new MqttContext().setUserName(null);
+ new MqttContext().setPassword(null);
}
@Test
public void retainedPublishStoresMessage() {
- new Publish().publish(mock(ChannelHandlerContext.class),
publishMessage(RETAINED_TOPIC, "hello", true));
+ new Publish().publish(connectedContext(),
publishMessage(RETAINED_TOPIC, "hello", true));
await().atMost(Duration.ofSeconds(5))
.until(() ->
"hello".equals(topicRepository.get(RETAINED_TOPIC)));
}
@Test
public void nonRetainedPublishDoesNotStoreMessage() {
- new Publish().publish(mock(ChannelHandlerContext.class),
publishMessage(NON_RETAINED_TOPIC, "hello", false));
+ new Publish().publish(connectedContext(),
publishMessage(NON_RETAINED_TOPIC, "hello", false));
assertNull(topicRepository.get(NON_RETAINED_TOPIC));
}
+ @Test
+ public void publishBeforeConnectClosesChannel() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+ new Publish().publish(ctx, publishMessage(UNCONNECTED_TOPIC, "hello",
true));
+
+ channel.runPendingTasks();
+ assertFalse(channel.isActive());
+ assertNull(topicRepository.get(UNCONNECTED_TOPIC));
+ }
+
+ @Test
+ public void publishAfterConnectOnSameChannelIsAccepted() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+ new Connect().connect(ctx, connectMessage());
+ new Publish().publish(ctx, publishMessage(END_TO_END_TOPIC, "hello",
true));
+
+ await().atMost(Duration.ofSeconds(5))
+ .until(() ->
"hello".equals(topicRepository.get(END_TO_END_TOPIC)));
+ }
+
@Test
public void zeroByteRetainedPublishClearsRetainedMessage() {
Publish publish = new Publish();
- publish.publish(mock(ChannelHandlerContext.class),
publishMessage(CLEARED_TOPIC, "hello", true));
+ publish.publish(connectedContext(), publishMessage(CLEARED_TOPIC,
"hello", true));
await().atMost(Duration.ofSeconds(5))
.until(() ->
"hello".equals(topicRepository.get(CLEARED_TOPIC)));
- publish.publish(mock(ChannelHandlerContext.class),
publishMessage(CLEARED_TOPIC, "", true));
+ publish.publish(connectedContext(), publishMessage(CLEARED_TOPIC, "",
true));
assertNull(topicRepository.get(CLEARED_TOPIC));
}
+ private ChannelHandlerContext connectedContext() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ new MessageType().setConnected(channel, true);
+ return channel.pipeline().lastContext();
+ }
+
+ private MqttConnectMessage connectMessage() {
+ MqttFixedHeader fixedHeader = new
MqttFixedHeader(MqttMessageType.CONNECT, false, MqttQoS.AT_MOST_ONCE, false, 0);
+ MqttConnectVariableHeader variableHeader = new
MqttConnectVariableHeader(
+ MqttVersion.MQTT_3_1_1.protocolName(),
MqttVersion.MQTT_3_1_1.protocolLevel(),
+ true, true, false, 0, false, false, 60);
+ MqttConnectPayload payload = new MqttConnectPayload(CLIENT_ID, null,
null,
+ USER_NAME, PASSWORD.getBytes(StandardCharsets.UTF_8));
+ return new MqttConnectMessage(fixedHeader, variableHeader, payload);
+ }
+
private MqttPublishMessage publishMessage(final String topic, final String
payload, final boolean retain) {
MqttFixedHeader fixedHeader = new
MqttFixedHeader(MqttMessageType.PUBLISH, false, MqttQoS.AT_MOST_ONCE, retain,
0);
MqttPublishVariableHeader variableHeader = new
MqttPublishVariableHeader(topic, 1);
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/UnsubscribeTest.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/UnsubscribeTest.java
new file mode 100644
index 0000000000..6d942b0b65
--- /dev/null
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/UnsubscribeTest.java
@@ -0,0 +1,58 @@
+/*
+ * 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.shenyu.protocol.mqtt;
+
+import io.netty.channel.ChannelHandlerContext;
+import io.netty.channel.ChannelInboundHandlerAdapter;
+import io.netty.channel.embedded.EmbeddedChannel;
+import io.netty.handler.codec.mqtt.MqttFixedHeader;
+import io.netty.handler.codec.mqtt.MqttMessageIdVariableHeader;
+import io.netty.handler.codec.mqtt.MqttMessageType;
+import io.netty.handler.codec.mqtt.MqttQoS;
+import io.netty.handler.codec.mqtt.MqttUnsubscribeMessage;
+import io.netty.handler.codec.mqtt.MqttUnsubscribePayload;
+import org.junit.jupiter.api.Test;
+
+import java.util.Collections;
+
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertNull;
+
+/**
+ * Test cases for {@link Unsubscribe}.
+ */
+public final class UnsubscribeTest {
+
+ @Test
+ public void unsubscribeBeforeConnectClosesChannel() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+ new Unsubscribe().unsubscribe(ctx, unsubscribeMessage());
+
+ channel.runPendingTasks();
+ assertFalse(channel.isActive());
+ assertNull(channel.readOutbound());
+ }
+
+ private MqttUnsubscribeMessage unsubscribeMessage() {
+ MqttFixedHeader fixedHeader = new
MqttFixedHeader(MqttMessageType.UNSUBSCRIBE, false, MqttQoS.AT_MOST_ONCE,
false, 0);
+ return new MqttUnsubscribeMessage(fixedHeader,
MqttMessageIdVariableHeader.from(1),
+ new
MqttUnsubscribePayload(Collections.singletonList("test/topic")));
+ }
+}