diff --git a/shenyu-protocol/shenyu-protocol-mqtt/pom.xml b/shenyu-protocol/shenyu-protocol-mqtt/pom.xml
index cf8e31fb56de..6ce05fbb8b98 100644
--- a/shenyu-protocol/shenyu-protocol-mqtt/pom.xml
+++ b/shenyu-protocol/shenyu-protocol-mqtt/pom.xml
@@ -52,6 +52,18 @@
junit-jupitertest
+
+ org.mockito
+ mockito-junit-jupiter
+ ${mockito.version}
+ test
+
+
+ org.mockito
+ mockito-core
+ ${mockito.version}
+ test
+
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 b23b07a661a3..fe7f7154b78f 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
@@ -26,6 +26,7 @@
import org.apache.commons.lang3.StringUtils;
import org.apache.shenyu.common.utils.Singleton;
import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
@@ -74,6 +75,17 @@ public void connect(final ChannelHandlerContext ctx, final MqttConnectMessage ms
// record connect
Singleton.INST.get(ChannelRepository.class).add(ctx.channel(), clientId);
+
+ // store will if present
+ if (msg.variableHeader().isWillFlag()) {
+ WillRepository.WillEntry will = new WillRepository.WillEntry(
+ msg.payload().willTopic(),
+ msg.payload().willMessageInBytes(),
+ msg.variableHeader().willQos(),
+ msg.variableHeader().isWillRetain());
+ Singleton.INST.get(WillRepository.class).add(ctx.channel(), will);
+ }
+
MqttConnAckMessage ackMessage = MqttMessageBuilders.connAck()
.returnCode(MqttConnectReturnCode.CONNECTION_ACCEPTED)
.sessionPresent(true)
diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Disconnect.java b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Disconnect.java
index d802247732be..4bb1bc04c7f2 100644
--- a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Disconnect.java
+++ b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Disconnect.java
@@ -22,6 +22,7 @@
import org.apache.shenyu.common.utils.Singleton;
import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
import org.apache.shenyu.protocol.mqtt.utils.MqttPacketIdGenerator;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository;
/**
* The DISCONNECT message is sent from the client to the server to indicate
@@ -37,8 +38,7 @@ public class Disconnect extends MessageType {
@Override
public void disconnect(final ChannelHandlerContext ctx) {
- //// todo Last words
- //// todo Clean session
+ Singleton.INST.get(WillRepository.class).remove(ctx.channel());
cleanChannel(ctx.channel());
ctx.close();
}
diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttFactory.java b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttFactory.java
index 59954045b573..049c24e50602 100644
--- a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttFactory.java
+++ b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttFactory.java
@@ -68,8 +68,10 @@ public void connect() {
case PUBREL:
messageType.pubRel(ctx, msg);
break;
- case PUBACK:
case DISCONNECT:
+ messageType.disconnect(ctx);
+ break;
+ case PUBACK:
default:
break;
}
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 c1d11817471b..1141652f10ce 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
@@ -29,6 +29,9 @@
import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
import org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository;
import org.apache.shenyu.protocol.mqtt.utils.MqttPacketIdGenerator;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository;
+
+import java.util.Objects;
/**
* mqtt transport handler.
@@ -51,8 +54,19 @@ public void channelRead(final ChannelHandlerContext ctx, final Object msg) throw
@Override
public void channelInactive(final ChannelHandlerContext ctx) throws Exception {
- Singleton.INST.get(ChannelRepository.class).remove(ctx.channel());
- ctx.fireChannelInactive();
+ final Channel channel = ctx.channel();
+ Singleton.INST.get(ChannelRepository.class).remove(channel);
+
+ final WillRepository willRepository = Singleton.INST.get(WillRepository.class);
+ final WillRepository.WillEntry will = willRepository.get(channel);
+ if (Objects.nonNull(will)) {
+ // a will is published at most once, and the repository keeps a strong reference
+ // to the channel, so it must be removed even if publishing fails.
+ willRepository.remove(channel);
+ Publish.publishWill(will);
+ }
+ // local state is consistent now, notify the rest of the pipeline exactly once.
+ super.channelInactive(ctx);
}
@Override
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 279b1e5c0403..56acd1cdc24a 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,9 @@
import org.apache.shenyu.protocol.mqtt.repositories.TopicRepository;
import org.apache.shenyu.protocol.mqtt.utils.MqttPacketIdGenerator;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository;
+
+import java.util.Objects;
import java.util.Map;
import java.util.concurrent.CompletableFuture;
@@ -143,4 +146,31 @@ private void send(final String topic, final ByteBuf payload, final MqttQoS publi
private static MqttQoS minQoS(final MqttQoS publishQoS, final MqttQoS grantedQoS) {
return publishQoS.value() <= grantedQoS.value() ? publishQoS : grantedQoS;
}
+
+ /**
+ * Publish a Last Will message to all subscribers of the will topic.
+ *
+ * @param will the will entry containing topic, message, qos, and retain flag
+ */
+ static void publishWill(final WillRepository.WillEntry will) {
+ if (Objects.isNull(will) || Objects.isNull(will.getTopic()) || Objects.isNull(will.getMessage())) {
+ return;
+ }
+ final Map subscribers = Singleton.INST.get(SubscribeRepository.class).getChannelsByTopic(will.getTopic());
+ final MqttQoS willQos = MqttQoS.valueOf(will.getQos());
+ subscribers.entrySet().parallelStream().forEach(entry -> {
+ Channel channel = entry.getKey();
+ if (channel.isActive()) {
+ MqttQoS qos = minQoS(willQos, entry.getValue());
+ int packetId = MqttQoS.AT_MOST_ONCE == qos
+ ? 0
+ : java.util.concurrent.ThreadLocalRandom.current().nextInt(1, 65536);
+ MqttFixedHeader mqttFixedHeader = new MqttFixedHeader(MqttMessageType.PUBLISH, false, qos, will.isRetain(), 0);
+ MqttPublishVariableHeader mqttPublishVariableHeader = new MqttPublishVariableHeader(will.getTopic(), packetId);
+ MqttPublishMessage mqttPublishMessage = new MqttPublishMessage(mqttFixedHeader, mqttPublishVariableHeader,
+ Unpooled.wrappedBuffer(will.getMessage()));
+ channel.writeAndFlush(mqttPublishMessage);
+ }
+ });
+ }
}
diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/TopicMatcher.java b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/TopicMatcher.java
new file mode 100644
index 000000000000..850b3addd314
--- /dev/null
+++ b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/TopicMatcher.java
@@ -0,0 +1,103 @@
+/*
+ * 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 java.util.Objects;
+
+/**
+ * MQTT topic filter matching per MQTT-4.7.
+ *
+ *
+ matches exactly one topic level.
+ *
+ *
# matches any number of subsequent levels (must appear at the end of the filter).
+ */
+public final class TopicMatcher {
+
+ private TopicMatcher() {
+ }
+
+ /**
+ * Check whether a topic filter matches a topic name.
+ *
+ * @param filter the subscription topic filter (may contain + and # wildcards)
+ * @param topic the published topic name (no wildcards)
+ * @return true if the filter matches the topic
+ */
+ public static boolean matches(final String filter, final String topic) {
+ if (Objects.isNull(filter) || Objects.isNull(topic)) {
+ return false;
+ }
+
+ // $ topics must not be matched by wildcards at the first level
+ if (topic.startsWith("$") && filter.length() > 0 && (filter.charAt(0) == '+' || filter.charAt(0) == '#')) {
+ return false;
+ }
+
+ String[] filterLevels = filter.split("/", -1);
+ String[] topicLevels = topic.split("/", -1);
+
+ int filterLen = filterLevels.length;
+ int topicLen = topicLevels.length;
+
+ for (int i = 0; i < filterLen; i++) {
+ String f = filterLevels[i];
+
+ if ("#".equals(f)) {
+ // MQTT-4.7.1-2: # matches any number of levels including the parent level
+ return i == filterLen - 1;
+ }
+
+ if (i >= topicLen) {
+ return false;
+ }
+
+ if (!"+".equals(f) && !f.equals(topicLevels[i])) {
+ return false;
+ }
+ }
+
+ return filterLen == topicLen;
+ }
+
+ /**
+ * Validate a topic filter per MQTT-4.7.1: wildcards must occupy an entire
+ * level, and # must be the last level. Filters must not be empty or
+ * contain the null character.
+ *
+ * @param filter the subscription topic filter
+ * @return true if the filter is valid
+ */
+ public static boolean isValidFilter(final String filter) {
+ if (Objects.isNull(filter) || filter.isEmpty() || filter.indexOf((char) 0) >= 0) {
+ return false;
+ }
+ String[] levels = filter.split("/", -1);
+ for (int i = 0; i < levels.length; i++) {
+ String level = levels[i];
+ if (level.indexOf('+') >= 0 || level.indexOf('#') >= 0) {
+ if (level.length() > 1) {
+ return false;
+ }
+ if ("#".equals(level) && i != levels.length - 1) {
+ return false;
+ }
+ }
+ }
+ return true;
+ }
+}
diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepository.java b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepository.java
index befd4a5c8e9b..766340b3604f 100644
--- a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepository.java
+++ b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepository.java
@@ -20,6 +20,7 @@
import io.netty.channel.Channel;
import io.netty.handler.codec.mqtt.MqttQoS;
import io.netty.handler.codec.mqtt.MqttTopicSubscription;
+import org.apache.shenyu.protocol.mqtt.TopicMatcher;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
@@ -109,4 +110,35 @@ private static MqttQoS maxQoS(final MqttQoS qos1, final MqttQoS qos2) {
return qos1.value() >= qos2.value() ? qos1 : qos2;
}
+ /**
+ * Get the channels whose subscription filter matches the published topic,
+ * mapped to the maximum qos granted across all their matching filters.
+ * Supports MQTT wildcards: + (single-level) and # (multi-level).
+ *
+ * @param topic the published topic name
+ * @return matching channels with their maximum granted qos
+ */
+ public Map getChannelsByTopic(final String topic) {
+ // MQTT requires at most one delivery per publish per client, so merge the
+ // granted qos when overlapping filters (e.g. sport/# and #) both match.
+ Map result = new ConcurrentHashMap<>();
+
+ // fast path: exact subscription, no wildcard scan needed
+ Map exactMatch = TOPIC_CHANNEL_FACTORY.get(topic);
+ if (Objects.nonNull(exactMatch)) {
+ result.putAll(exactMatch);
+ }
+
+ for (Map.Entry> entry : TOPIC_CHANNEL_FACTORY.entrySet()) {
+ String filter = entry.getKey();
+ if (filter.equals(topic) || filter.indexOf('+') < 0 && filter.indexOf('#') < 0) {
+ continue;
+ }
+ if (TopicMatcher.matches(filter, topic)) {
+ entry.getValue().forEach((channel, qos) -> result.merge(channel, qos, SubscribeRepository::maxQoS));
+ }
+ }
+ return result;
+ }
+
}
diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/WillRepository.java b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/WillRepository.java
new file mode 100644
index 000000000000..9a57439ae470
--- /dev/null
+++ b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/WillRepository.java
@@ -0,0 +1,85 @@
+/*
+ * 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.repositories;
+
+import io.netty.channel.Channel;
+
+import java.util.Map;
+import java.util.concurrent.ConcurrentHashMap;
+
+/**
+ * Stores Last Will and Testament for connected clients.
+ * Will is set on CONNECT and cleared on graceful DISCONNECT.
+ * On ungraceful disconnect (channelInactive with will present), the will is published.
+ */
+public class WillRepository implements BaseRepository {
+
+ private static final Map WILL_FACTORY = new ConcurrentHashMap<>();
+
+ @Override
+ public void add(final Channel channel, final WillEntry willEntry) {
+ WILL_FACTORY.put(channel, willEntry);
+ }
+
+ @Override
+ public void remove(final Channel channel) {
+ WILL_FACTORY.remove(channel);
+ }
+
+ @Override
+ public WillEntry get(final Channel channel) {
+ return WILL_FACTORY.get(channel);
+ }
+
+ /**
+ * Holds the will message fields from a CONNECT payload.
+ */
+ public static class WillEntry {
+
+ private final String topic;
+
+ private final byte[] message;
+
+ private final int qos;
+
+ private final boolean retain;
+
+ public WillEntry(final String topic, final byte[] message, final int qos, final boolean retain) {
+ this.topic = topic;
+ this.message = message;
+ this.qos = qos;
+ this.retain = retain;
+ }
+
+ public String getTopic() {
+ return topic;
+ }
+
+ public byte[] getMessage() {
+ return message;
+ }
+
+ public int getQos() {
+ return qos;
+ }
+
+ public boolean isRetain() {
+ return retain;
+ }
+ }
+}
\ No newline at end of file
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 2b8acc8e33d5..011b7598921b 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
@@ -30,16 +30,23 @@
import io.netty.handler.codec.mqtt.MqttVersion;
import org.apache.shenyu.common.utils.Singleton;
import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository.WillEntry;
import org.junit.jupiter.api.AfterAll;
+import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.BeforeAll;
+import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import java.nio.charset.StandardCharsets;
-import java.time.Duration;
+import java.util.ArrayList;
+import java.util.List;
import static io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_ACCEPTED;
+import static io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_REFUSED_BAD_USER_NAME_OR_PASSWORD;
+import static io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_REFUSED_IDENTIFIER_REJECTED;
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.assertArrayEquals;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNotNull;
@@ -51,26 +58,55 @@
*/
public final class ConnectTest {
- 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 ChannelRepository channelRepository;
+ private static final String CLIENT_ID = "test-client";
+
+ private static final String WILL_TOPIC = "status/client-001";
+
+ private final List channels = new ArrayList<>();
+
+ private ChannelRepository channelRepository;
+
+ private WillRepository willRepository;
+
+ private Connect connect;
@BeforeAll
- static void setUp() {
+ static void setUpCredentials() {
+ MqttContext mqttContext = new MqttContext();
+ mqttContext.setUserName(USER_NAME);
+ mqttContext.setPassword(PASSWORD);
+ }
+
+ @AfterAll
+ static void clearCredentials() {
+ MqttContext mqttContext = new MqttContext();
+ mqttContext.setUserName(null);
+ mqttContext.setPassword(null);
+ }
+
+ @BeforeEach
+ public void setUp() {
+ connect = new Connect();
channelRepository = new ChannelRepository();
+ willRepository = new WillRepository();
Singleton.INST.single(ChannelRepository.class, channelRepository);
- new MqttContext().setUserName(USER_NAME);
- new MqttContext().setPassword(PASSWORD);
+ Singleton.INST.single(WillRepository.class, willRepository);
}
- @AfterAll
- static void tearDown() {
- new MqttContext().setUserName(null);
- new MqttContext().setPassword(null);
+ @AfterEach
+ public void tearDown() {
+ for (EmbeddedChannel channel : channels) {
+ channelRepository.remove(channel);
+ willRepository.remove(channel);
+ channel.finishAndReleaseAll();
+ }
+ channels.clear();
+ Singleton.INST.single(ChannelRepository.class, new ChannelRepository());
+ Singleton.INST.single(WillRepository.class, new WillRepository());
}
@Test
@@ -90,55 +126,169 @@ public void mqtt5ConnectIsAccepted() {
@Test
public void unsupportedProtocolVersionIsRejected() {
- EmbeddedChannel channel = new EmbeddedChannel(new ChannelInboundHandlerAdapter());
- ChannelHandlerContext ctx = channel.pipeline().lastContext();
+ EmbeddedChannel channel = newChannel();
- new Connect().connect(ctx, connectMessage("MQTT", 6));
+ connect.connect(context(channel), connectMessage("MQTT", 6));
MqttConnAckMessage ackMessage = channel.readOutbound();
assertNotNull(ackMessage);
- assertEquals(CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION,
- ackMessage.variableHeader().connectReturnCode());
- channel.runPendingTasks();
- assertFalse(channel.isActive());
+ assertEquals(CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION, ackMessage.variableHeader().connectReturnCode());
+ assertFalse(ackMessage.variableHeader().isSessionPresent());
+ assertChannelClosed(channel);
+ assertNull(channelRepository.get(channel));
+ }
+
+ @Test
+ public void emptyClientIdIsRejected() {
+ EmbeddedChannel channel = newChannel();
+
+ connect.connect(context(channel), buildConnectMessage(MqttVersion.MQTT_3_1_1.protocolName(),
+ MqttVersion.MQTT_3_1_1.protocolLevel(), "", PASSWORD, false, 0, false, null, null));
+
+ MqttConnAckMessage ackMessage = channel.readOutbound();
+ assertNotNull(ackMessage);
+ assertEquals(CONNECTION_REFUSED_IDENTIFIER_REJECTED, ackMessage.variableHeader().connectReturnCode());
+ assertChannelClosed(channel);
+ assertNull(channelRepository.get(channel));
+ }
+
+ @Test
+ public void invalidCredentialsAreRejected() {
+ EmbeddedChannel channel = newChannel();
+
+ connect.connect(context(channel), buildConnectMessage(MqttVersion.MQTT_3_1_1.protocolName(),
+ MqttVersion.MQTT_3_1_1.protocolLevel(), CLIENT_ID, "invalid-password", false, 0, false, null, null));
+
+ MqttConnAckMessage ackMessage = channel.readOutbound();
+ assertNotNull(ackMessage);
+ assertEquals(CONNECTION_REFUSED_BAD_USER_NAME_OR_PASSWORD, ackMessage.variableHeader().connectReturnCode());
+ assertChannelClosed(channel);
assertNull(channelRepository.get(channel));
}
@Test
public void duplicateConnectIsRejected() {
- EmbeddedChannel channel = new EmbeddedChannel(new ChannelInboundHandlerAdapter());
- ChannelHandlerContext ctx = channel.pipeline().lastContext();
+ EmbeddedChannel channel = newChannel();
+ ChannelHandlerContext ctx = context(channel);
- new Connect().connect(ctx, connectMessage(MqttVersion.MQTT_3_1_1.protocolName(), MqttVersion.MQTT_3_1_1.protocolLevel()));
+ connect.connect(ctx, connectMessage(MqttVersion.MQTT_3_1_1));
assertNotNull(channel.readOutbound());
- new Connect().connect(ctx, connectMessage(MqttVersion.MQTT_3_1_1.protocolName(), MqttVersion.MQTT_3_1_1.protocolLevel()));
+ connect.connect(ctx, connectMessage(MqttVersion.MQTT_3_1_1));
- channel.runPendingTasks();
- assertFalse(channel.isActive());
+ assertChannelClosed(channel);
assertNull(channel.readOutbound());
}
+ @Test
+ public void testStoresWillOnConnect() {
+ EmbeddedChannel channel = newChannel();
+ byte[] willMessage = "client disconnected unexpectedly".getBytes(StandardCharsets.UTF_8);
+
+ connect.connect(context(channel), willConnectMessage(1, true, WILL_TOPIC, willMessage));
+
+ WillEntry will = willRepository.get(channel);
+ assertNotNull(will);
+ assertEquals(WILL_TOPIC, will.getTopic());
+ assertArrayEquals(willMessage, will.getMessage());
+ assertEquals(1, will.getQos());
+ assertTrue(will.isRetain());
+ }
+
+ @Test
+ public void testDoesNotStoreWillWhenWillFlagIsFalse() {
+ EmbeddedChannel channel = newChannel();
+
+ connect.connect(context(channel), connectMessage(MqttVersion.MQTT_3_1_1));
+
+ assertNull(willRepository.get(channel));
+ }
+
+ @Test
+ public void testWillQosZero() {
+ EmbeddedChannel channel = newChannel();
+ byte[] willMessage = "qos0 will".getBytes(StandardCharsets.UTF_8);
+
+ connect.connect(context(channel), willConnectMessage(0, false, "topic/qos0", willMessage));
+
+ WillEntry will = willRepository.get(channel);
+ assertNotNull(will);
+ assertEquals(0, will.getQos());
+ assertFalse(will.isRetain());
+ }
+
+ @Test
+ public void testWillRetainTrue() {
+ EmbeddedChannel channel = newChannel();
+ byte[] willMessage = "retained will".getBytes(StandardCharsets.UTF_8);
+
+ connect.connect(context(channel), willConnectMessage(2, true, "topic/retained", willMessage));
+
+ WillEntry will = willRepository.get(channel);
+ assertNotNull(will);
+ assertEquals(2, will.getQos());
+ assertTrue(will.isRetain());
+ }
+
+ @Test
+ public void willIsNotStoredWhenConnectIsRejected() {
+ EmbeddedChannel channel = newChannel();
+
+ connect.connect(context(channel), buildConnectMessage("MQTT", 6, CLIENT_ID, PASSWORD,
+ true, 1, true, WILL_TOPIC, "retained will".getBytes(StandardCharsets.UTF_8)));
+
+ assertNull(willRepository.get(channel));
+ }
+
private void connectIsAccepted(final MqttVersion version) {
- EmbeddedChannel channel = new EmbeddedChannel(new ChannelInboundHandlerAdapter());
- ChannelHandlerContext ctx = channel.pipeline().lastContext();
+ EmbeddedChannel channel = newChannel();
- new Connect().connect(ctx, connectMessage(version.protocolName(), version.protocolLevel()));
+ connect.connect(context(channel), connectMessage(version));
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)));
+ assertEquals(CLIENT_ID, channelRepository.get(channel));
+ }
+
+ private void assertChannelClosed(final EmbeddedChannel channel) {
+ channel.runPendingTasks();
+ assertFalse(channel.isActive());
+ }
+
+ private EmbeddedChannel newChannel() {
+ EmbeddedChannel channel = new EmbeddedChannel(new ChannelInboundHandlerAdapter());
+ channels.add(channel);
+ return channel;
+ }
+
+ private ChannelHandlerContext context(final EmbeddedChannel channel) {
+ return channel.pipeline().lastContext();
+ }
+
+ private MqttConnectMessage connectMessage(final MqttVersion version) {
+ return connectMessage(version.protocolName(), version.protocolLevel());
}
private MqttConnectMessage connectMessage(final String protocolName, final int protocolLevel) {
+ return buildConnectMessage(protocolName, protocolLevel, CLIENT_ID, PASSWORD, false, 0, false, null, null);
+ }
+
+ private MqttConnectMessage willConnectMessage(final int willQos, final boolean willRetain,
+ final String willTopic, final byte[] willMessage) {
+ return buildConnectMessage(MqttVersion.MQTT_3_1_1.protocolName(), MqttVersion.MQTT_3_1_1.protocolLevel(),
+ CLIENT_ID, PASSWORD, true, willQos, willRetain, willTopic, willMessage);
+ }
+
+ private MqttConnectMessage buildConnectMessage(final String protocolName, final int protocolLevel,
+ final String clientId, final String password, final boolean willFlag, final int willQos,
+ final boolean willRetain, final String willTopic, final byte[] willMessage) {
MqttFixedHeader fixedHeader = new MqttFixedHeader(MqttMessageType.CONNECT, false, MqttQoS.AT_MOST_ONCE, false, 0);
MqttConnectVariableHeader variableHeader = new MqttConnectVariableHeader(protocolName, protocolLevel,
- true, true, false, 0, false, false, 60);
- MqttConnectPayload payload = new MqttConnectPayload(CLIENT_ID, null, null,
- USER_NAME, PASSWORD.getBytes(StandardCharsets.UTF_8));
+ true, true, willRetain, willQos, willFlag, false, 60);
+ MqttConnectPayload payload = new MqttConnectPayload(clientId, willTopic, willMessage,
+ 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/DisconnectTest.java b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/DisconnectTest.java
new file mode 100644
index 000000000000..7e2006e7fef9
--- /dev/null
+++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/DisconnectTest.java
@@ -0,0 +1,89 @@
+/*
+ * 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.Channel;
+import io.netty.channel.ChannelHandlerContext;
+import org.apache.shenyu.common.utils.Singleton;
+import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository;
+import org.junit.jupiter.api.AfterEach;
+import org.junit.jupiter.api.BeforeEach;
+import org.junit.jupiter.api.Test;
+import org.junit.jupiter.api.extension.ExtendWith;
+import org.mockito.Mock;
+import org.mockito.junit.jupiter.MockitoExtension;
+
+import static org.hamcrest.MatcherAssert.assertThat;
+import static org.hamcrest.Matchers.nullValue;
+import static org.mockito.Mockito.when;
+
+@ExtendWith(MockitoExtension.class)
+public class DisconnectTest {
+
+ @Mock
+ private ChannelHandlerContext ctx;
+
+ @Mock
+ private Channel channel;
+
+ private Disconnect disconnect;
+
+ private WillRepository willRepository;
+
+ private ChannelRepository channelRepository;
+
+ @BeforeEach
+ public void setUp() {
+ disconnect = new Disconnect();
+ willRepository = new WillRepository();
+ channelRepository = new ChannelRepository();
+ Singleton.INST.single(WillRepository.class, willRepository);
+ Singleton.INST.single(ChannelRepository.class, channelRepository);
+ when(ctx.channel()).thenReturn(channel);
+ }
+
+ @AfterEach
+ public void tearDown() {
+ Singleton.INST.single(WillRepository.class, new WillRepository());
+ Singleton.INST.single(ChannelRepository.class, new ChannelRepository());
+ }
+
+ @Test
+ public void testDisconnectClearsWill() {
+ WillRepository.WillEntry will = new WillRepository.WillEntry("topic/will", "goodbye".getBytes(), 0, false);
+ willRepository.add(channel, will);
+
+ disconnect.disconnect(ctx);
+
+ assertThat(willRepository.get(channel), nullValue());
+ }
+
+ @Test
+ public void testDisconnectWithoutWill() {
+ disconnect.disconnect(ctx);
+ assertThat(willRepository.get(channel), nullValue());
+ }
+
+ @Test
+ public void testDisconnectRemovesChannel() {
+ channelRepository.add(channel, "client-1");
+ disconnect.disconnect(ctx);
+ assertThat(channelRepository.get(channel), nullValue());
+ }
+}
diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttFactoryTest.java b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttFactoryTest.java
index deb0550121a5..f92bf2f9fd99 100644
--- a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttFactoryTest.java
+++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttFactoryTest.java
@@ -40,6 +40,8 @@
import io.netty.handler.codec.mqtt.MqttVersion;
import org.apache.shenyu.common.utils.Singleton;
import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository.WillEntry;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
@@ -49,6 +51,7 @@
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;
/**
@@ -62,15 +65,24 @@ public final class MqttFactoryTest {
private static final String PASSWORD = "factory-password";
+ private ChannelRepository channelRepository;
+
+ private WillRepository willRepository;
+
@BeforeEach
public void setUp() {
- Singleton.INST.single(ChannelRepository.class, new ChannelRepository());
+ channelRepository = new ChannelRepository();
+ willRepository = new WillRepository();
+ Singleton.INST.single(ChannelRepository.class, channelRepository);
+ Singleton.INST.single(WillRepository.class, willRepository);
new MqttContext().setUserName(USER_NAME);
new MqttContext().setPassword(PASSWORD);
}
@AfterEach
public void tearDown() {
+ Singleton.INST.single(ChannelRepository.class, new ChannelRepository());
+ Singleton.INST.single(WillRepository.class, new WillRepository());
new MqttContext().setUserName(null);
new MqttContext().setPassword(null);
}
@@ -149,13 +161,18 @@ public void pubAckShouldFallThroughToNoOp() {
}
@Test
- public void disconnectShouldFallThroughToNoOp() {
+ public void disconnectShouldBeDispatchedAndClearConnectionState() {
EmbeddedChannel channel = new EmbeddedChannel(new ChannelInboundHandlerAdapter());
ChannelHandlerContext ctx = channel.pipeline().lastContext();
+ channelRepository.add(channel, CLIENT_ID);
+ willRepository.add(channel, new WillEntry("status/will", "goodbye".getBytes(StandardCharsets.UTF_8), 1, true));
new MqttFactory(new MqttMessage(fixedHeader(MqttMessageType.DISCONNECT)), ctx).connect();
+ channel.runPendingTasks();
- assertTrue(channel.isActive());
+ assertFalse(channel.isActive());
+ assertNull(channelRepository.get(channel));
+ assertNull(willRepository.get(channel));
channel.finishAndReleaseAll();
}
diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java
index 58c24665797a..d609b500ca5c 100644
--- a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java
+++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java
@@ -19,6 +19,9 @@
import io.netty.buffer.ByteBuf;
import io.netty.buffer.Unpooled;
+import io.netty.channel.Channel;
+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;
@@ -34,29 +37,41 @@
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.WillRepository;
import org.apache.shenyu.protocol.mqtt.utils.MqttPacketIdGenerator;
import org.awaitility.core.ThrowingRunnable;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
+import org.junit.jupiter.api.extension.ExtendWith;
+import org.mockito.ArgumentCaptor;
+import org.mockito.Mock;
+import org.mockito.junit.jupiter.MockitoExtension;
import java.nio.charset.StandardCharsets;
import java.time.Duration;
import java.util.Collections;
+import java.util.concurrent.atomic.AtomicInteger;
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.lenient;
+import static org.mockito.Mockito.verify;
+import static org.mockito.Mockito.when;
/**
* Test cases for {@link MqttTransportHandler}.
*/
+@ExtendWith(MockitoExtension.class)
public final class MqttTransportHandlerTest {
private static final String TOPIC = "test/topic";
+ private static final String WILL_TOPIC = "status/offline";
+
private static final String CLIENT_ID = "test-client";
private static final String USER_NAME = "test-user";
@@ -75,14 +90,29 @@ public final class MqttTransportHandlerTest {
private static final SubscribeRepository SUBSCRIBE_REPOSITORY = new SubscribeRepository();
+ @Mock
+ private ChannelHandlerContext ctx;
+
+ @Mock
+ private Channel channel;
+
+ private MqttTransportHandler handler;
+
+ private WillRepository willRepository;
+
private EmbeddedChannel registeredChannel;
@BeforeEach
public void setUp() {
+ handler = new MqttTransportHandler();
+ willRepository = new WillRepository();
+ Singleton.INST.single(WillRepository.class, willRepository);
Singleton.INST.single(ChannelRepository.class, CHANNEL_REPOSITORY);
Singleton.INST.single(SubscribeRepository.class, SUBSCRIBE_REPOSITORY);
new MqttContext().setUserName(USER_NAME);
new MqttContext().setPassword(PASSWORD);
+ // only the will tests drive the handler through a mocked context
+ lenient().when(ctx.channel()).thenReturn(channel);
registeredChannel = new EmbeddedChannel();
CHANNEL_REPOSITORY.add(registeredChannel, CLIENT_ID);
@@ -96,74 +126,79 @@ public void setUp() {
@AfterEach
public void tearDown() {
+ willRepository.remove(channel);
+ SUBSCRIBE_REPOSITORY.remove(channel);
MqttPacketIdGenerator.remove(registeredChannel);
CHANNEL_REPOSITORY.remove(registeredChannel);
SUBSCRIBE_REPOSITORY.remove(registeredChannel);
- awaitAssert(() -> assertFalse(SUBSCRIBE_REPOSITORY.get(TOPIC).containsKey(registeredChannel)));
+ awaitAssert(() -> {
+ assertFalse(SUBSCRIBE_REPOSITORY.get(TOPIC).containsKey(registeredChannel));
+ assertTrue(SUBSCRIBE_REPOSITORY.get(WILL_TOPIC).isEmpty());
+ });
registeredChannel.finishAndReleaseAll();
-
new MqttContext().setUserName(null);
new MqttContext().setPassword(null);
}
@Test
public void channelReadReleasesInboundPublishMessage() {
- EmbeddedChannel channel = new EmbeddedChannel(new MqttTransportHandler());
+ EmbeddedChannel publisherChannel = new EmbeddedChannel(new MqttTransportHandler());
MqttFixedHeader fixedHeader = new MqttFixedHeader(MqttMessageType.PUBLISH, false, MqttQoS.AT_MOST_ONCE, false, 0);
MqttPublishVariableHeader variableHeader = new MqttPublishVariableHeader(TOPIC, 1);
MqttPublishMessage message = new MqttPublishMessage(fixedHeader, variableHeader,
Unpooled.copiedBuffer("hello", CharsetUtil.UTF_8));
ByteBuf payload = message.payload();
- channel.writeInbound(message);
+ publisherChannel.writeInbound(message);
assertEquals(0, payload.refCnt());
- channel.finishAndReleaseAll();
+ publisherChannel.finishAndReleaseAll();
}
@Test
public void duplicateConnectCleansUpChannelRepository() {
- EmbeddedChannel channel = new EmbeddedChannel(new MqttTransportHandler());
+ EmbeddedChannel sessionChannel = new EmbeddedChannel(new MqttTransportHandler());
- channel.writeInbound(connectMessage());
- assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(channel));
+ sessionChannel.writeInbound(connectMessage());
+ assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(sessionChannel));
- channel.writeInbound(connectMessage());
- channel.runPendingTasks();
+ sessionChannel.writeInbound(connectMessage());
+ sessionChannel.runPendingTasks();
- assertFalse(channel.isActive());
- assertNull(CHANNEL_REPOSITORY.get(channel));
- channel.finishAndReleaseAll();
+ assertFalse(sessionChannel.isActive());
+ assertNull(CHANNEL_REPOSITORY.get(sessionChannel));
+ sessionChannel.finishAndReleaseAll();
}
@Test
public void abruptChannelCloseCleansUpChannelRepository() {
- EmbeddedChannel channel = new EmbeddedChannel(new MqttTransportHandler());
+ EmbeddedChannel sessionChannel = new EmbeddedChannel(new MqttTransportHandler());
- channel.writeInbound(connectMessage());
- assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(channel));
+ sessionChannel.writeInbound(connectMessage());
+ assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(sessionChannel));
- channel.close();
- channel.runPendingTasks();
+ sessionChannel.close();
+ sessionChannel.runPendingTasks();
- assertFalse(channel.isActive());
- assertNull(CHANNEL_REPOSITORY.get(channel));
- channel.finishAndReleaseAll();
+ assertFalse(sessionChannel.isActive());
+ assertNull(CHANNEL_REPOSITORY.get(sessionChannel));
+ sessionChannel.finishAndReleaseAll();
}
@Test
public void nonMqttMessageClosesConnectedChannel() {
- EmbeddedChannel channel = new EmbeddedChannel(new MqttTransportHandler());
+ EmbeddedChannel sessionChannel = new EmbeddedChannel(new MqttTransportHandler());
- channel.writeInbound(connectMessage());
- assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(channel));
+ sessionChannel.writeInbound(connectMessage());
+ assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(sessionChannel));
- channel.writeInbound("not-a-mqtt-message");
+ sessionChannel.writeInbound("not-a-mqtt-message");
+ sessionChannel.runPendingTasks();
- assertFalse(channel.isActive());
- assertNull(CHANNEL_REPOSITORY.get(channel));
- channel.finishAndReleaseAll();
+ assertFalse(sessionChannel.isActive());
+ assertNull(CHANNEL_REPOSITORY.get(sessionChannel));
+ sessionChannel.finishAndReleaseAll();
}
@Test
@@ -177,6 +212,60 @@ public void testOperationCompleteCleansRepositoriesOnClose() throws Exception {
assertEquals(1, MqttPacketIdGenerator.next(registeredChannel));
}
+ @Test
+ public void channelInactiveIsPropagatedOnlyOnce() {
+ AtomicInteger fired = new AtomicInteger();
+ EmbeddedChannel sessionChannel = new EmbeddedChannel(new MqttTransportHandler(), new ChannelInboundHandlerAdapter() {
+ @Override
+ public void channelInactive(final ChannelHandlerContext context) throws Exception {
+ fired.incrementAndGet();
+ super.channelInactive(context);
+ }
+ });
+
+ sessionChannel.close();
+ sessionChannel.runPendingTasks();
+
+ assertEquals(1, fired.get());
+ sessionChannel.finishAndReleaseAll();
+ }
+
+ @Test
+ public void testChannelInactiveFiresWillAndRemovesIt() throws Exception {
+ willRepository.add(channel, new WillRepository.WillEntry(WILL_TOPIC, "sudden disconnect".getBytes(), 1, true));
+ SUBSCRIBE_REPOSITORY.add(channel,
+ Collections.singletonList(new MqttTopicSubscription(WILL_TOPIC, MqttQoS.AT_LEAST_ONCE)));
+ awaitAssert(() -> assertTrue(SUBSCRIBE_REPOSITORY.get(WILL_TOPIC).containsKey(channel)));
+ when(channel.isActive()).thenReturn(true);
+
+ handler.channelInactive(ctx);
+
+ ArgumentCaptor published = ArgumentCaptor.forClass(MqttPublishMessage.class);
+ verify(channel).writeAndFlush(published.capture());
+ assertEquals(WILL_TOPIC, published.getValue().variableHeader().topicName());
+ assertEquals(MqttQoS.AT_LEAST_ONCE, published.getValue().fixedHeader().qosLevel());
+ assertTrue(published.getValue().fixedHeader().isRetain());
+ assertNull(willRepository.get(channel));
+ }
+
+ @Test
+ public void testChannelInactiveDoesNothingWhenNoWill() throws Exception {
+ handler.channelInactive(ctx);
+
+ assertNull(willRepository.get(channel));
+ }
+
+ @Test
+ public void testChannelInactiveAfterDisconnectClearsWill() throws Exception {
+ willRepository.add(channel, new WillRepository.WillEntry("status/clean", "graceful close".getBytes(), 0, false));
+
+ // a graceful disconnect removes the will first, so channelInactive must not publish it
+ willRepository.remove(channel);
+ handler.channelInactive(ctx);
+
+ assertNull(willRepository.get(channel));
+ }
+
private MqttConnectMessage connectMessage() {
MqttFixedHeader fixedHeader = new MqttFixedHeader(MqttMessageType.CONNECT, false, MqttQoS.AT_MOST_ONCE, false, 0);
MqttConnectVariableHeader variableHeader = new MqttConnectVariableHeader(
diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishWillTest.java b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishWillTest.java
new file mode 100644
index 000000000000..4e89d1e8e63d
--- /dev/null
+++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishWillTest.java
@@ -0,0 +1,161 @@
+/*
+ * 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.Channel;
+import io.netty.handler.codec.mqtt.MqttPublishMessage;
+import io.netty.handler.codec.mqtt.MqttQoS;
+import org.apache.shenyu.common.utils.Singleton;
+import org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository;
+import org.apache.shenyu.protocol.mqtt.repositories.WillRepository;
+import org.junit.jupiter.api.AfterEach;
+import org.junit.jupiter.api.BeforeEach;
+import org.junit.jupiter.api.Test;
+import org.junit.jupiter.api.extension.ExtendWith;
+import org.mockito.ArgumentCaptor;
+import org.mockito.Mock;
+import org.mockito.junit.jupiter.MockitoExtension;
+
+import java.util.Collections;
+import java.util.HashMap;
+import java.util.Map;
+
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertTrue;
+import static org.mockito.ArgumentMatchers.any;
+import static org.mockito.Mockito.never;
+import static org.mockito.Mockito.verify;
+import static org.mockito.Mockito.when;
+
+@ExtendWith(MockitoExtension.class)
+public class PublishWillTest {
+
+ @Mock
+ private SubscribeRepository subscribeRepository;
+
+ @Mock
+ private Channel subscriberChannel;
+
+ @Mock
+ private Channel otherChannel;
+
+ @BeforeEach
+ public void setUp() {
+ Singleton.INST.single(SubscribeRepository.class, subscribeRepository);
+ }
+
+ @AfterEach
+ public void tearDown() {
+ Singleton.INST.single(SubscribeRepository.class, new SubscribeRepository());
+ }
+
+ @Test
+ public void testPublishWillToActiveSubscriber() {
+ when(subscriberChannel.isActive()).thenReturn(true);
+ when(subscribeRepository.getChannelsByTopic("status/offline"))
+ .thenReturn(Collections.singletonMap(subscriberChannel, MqttQoS.EXACTLY_ONCE));
+
+ byte[] message = "client lost".getBytes();
+ WillRepository.WillEntry will = new WillRepository.WillEntry("status/offline", message, 1, true);
+ Publish.publishWill(will);
+
+ ArgumentCaptor