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 b3f21ecf6e fix: deliver mqtt messages at granted qos with
per-subscriber packet id (#6913)
b3f21ecf6e is described below
commit b3f21ecf6e77ac28dc634ba8c931a8875eb16f3a
Author: wy471x <[email protected]>
AuthorDate: Wed Sep 30 10:00:23 2026 +0800
fix: deliver mqtt messages at granted qos with per-subscriber packet id
(#6913)
* fix: deliver mqtt messages at granted qos with per-subscriber packet id
Store the granted qos per subscriber channel in SubscribeRepository and
allocate a packet id from each subscriber's own id space when fanning out
publishes, instead of hard-coding AT_MOST_ONCE and reusing the publisher's
packet id for every subscriber.
Co-Authored-By: Claude Opus 4.7 <[email protected]>
* fix: prevent mqtt packet id space leak and payload over-release on fan-out
Use a weak-keyed map for per-channel packet id spaces so channels closed
without a DISCONNECT do not leak entries, and retain the publish payload
per subscriber so each outbound message owns a reference.
Co-Authored-By: Claude Opus 4.7 <[email protected]>
* fix: correct mqtt qos responses and complete qos2 handshake (#6743)
Respond to qos2 publishes with PUBREC instead of PUBACK, implement the
PUBREL -> PUBCOMP half-handshake, and zero the reserved fixed-header
bits of the PUBACK/PUBREC responses. Remove the unused qos0 stub.
Co-Authored-By: Claude Opus 4.7 <[email protected]>
* fix: clean up channel repositories when connection closes (#6744)
Remove the closing channel from ChannelRepository, SubscribeRepository
and the packet id space in MqttTransportHandler.operationComplete so
ungraceful disconnects do not leak entries.
Co-Authored-By: Claude Opus 4.7 <[email protected]>
* fix(mqtt): drop dead code from merge and isolate repository test state
- SubscribeRepository.remove(List, Channel) kept two blocks after the merge,
one of which assigns a Map to a List and does not compile
- release the static repository maps around every test in PublishTest,
SubscribeRepositoryTest and MqttTransportHandlerTest so the shared state
of one test class cannot leak into another
- cover the non-MQTT read path of MqttTransportHandler, the per-topic QoS
merging of SubscribeRepository and the retained publish handling
---------
Co-authored-by: Claude Opus 4.7 <[email protected]>
Co-authored-by: aias00 <[email protected]>
---
.../shenyu/protocol/mqtt/AbstractMessageType.java | 10 +
.../apache/shenyu/protocol/mqtt/Disconnect.java | 2 +
.../apache/shenyu/protocol/mqtt/MessageType.java | 7 +
.../apache/shenyu/protocol/mqtt/MqttFactory.java | 3 +
.../shenyu/protocol/mqtt/MqttTransportHandler.java | 9 +-
.../org/apache/shenyu/protocol/mqtt/PubRel.java | 41 +++
.../org/apache/shenyu/protocol/mqtt/Publish.java | 44 +--
.../mqtt/repositories/SubscribeRepository.java | 68 +++--
.../protocol/mqtt/utils/MqttPacketIdGenerator.java | 60 ++++
.../protocol/mqtt/MqttTransportHandlerTest.java | 98 ++++++-
.../apache/shenyu/protocol/mqtt/PubRelTest.java | 55 ++++
.../apache/shenyu/protocol/mqtt/PublishTest.java | 314 ++++++++++++++++++---
.../mqtt/repositories/SubscribeRepositoryTest.java | 190 +++++++++++--
.../mqtt/utils/MqttPacketIdGeneratorTest.java | 77 +++++
14 files changed, 863 insertions(+), 115 deletions(-)
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/AbstractMessageType.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/AbstractMessageType.java
index d541a68321..ffe02da978 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/AbstractMessageType.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/AbstractMessageType.java
@@ -19,6 +19,7 @@ package org.apache.shenyu.protocol.mqtt;
import io.netty.channel.ChannelHandlerContext;
import io.netty.handler.codec.mqtt.MqttConnectMessage;
+import io.netty.handler.codec.mqtt.MqttMessage;
import io.netty.handler.codec.mqtt.MqttPublishMessage;
import io.netty.handler.codec.mqtt.MqttSubscribeMessage;
import io.netty.handler.codec.mqtt.MqttUnsubscribeMessage;
@@ -88,4 +89,13 @@ public interface AbstractMessageType {
}
+ /**
+ * Publish Release, third message of the QoS 2 protocol flow.
+ * @param ctx ctx
+ * @param msg msg
+ */
+ default void pubRel(final ChannelHandlerContext ctx, final MqttMessage
msg) {
+
+ }
+
}
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 87a06e39d2..d802247732 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
@@ -21,6 +21,7 @@ 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.utils.MqttPacketIdGenerator;
/**
* The DISCONNECT message is sent from the client to the server to indicate
@@ -45,5 +46,6 @@ public class Disconnect extends MessageType {
private void cleanChannel(final Channel channel) {
//// todo ttl
Singleton.INST.get(ChannelRepository.class).remove(channel);
+ MqttPacketIdGenerator.remove(channel);
}
}
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 f6bdecac39..487405fe50 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
@@ -20,6 +20,7 @@ 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.MqttMessage;
import io.netty.handler.codec.mqtt.MqttPublishMessage;
import io.netty.handler.codec.mqtt.MqttSubscribeMessage;
import io.netty.handler.codec.mqtt.MqttUnsubscribeMessage;
@@ -91,4 +92,10 @@ public class MessageType implements AbstractMessageType {
//// todo polymorphism disconnect
new Disconnect().disconnect(ctx);
}
+
+ @Override
+ public void pubRel(final ChannelHandlerContext ctx, final MqttMessage msg)
{
+ //// todo polymorphism pubRel
+ new PubRel().pubRel(ctx, msg);
+ }
}
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 9a59f1407e..59954045b5 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
@@ -65,6 +65,9 @@ public class MqttFactory {
case PINGREQ:
messageType.pingReq(ctx);
break;
+ case PUBREL:
+ messageType.pubRel(ctx, msg);
+ break;
case PUBACK:
case DISCONNECT:
default:
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 a6fa975297..bc1ccad71f 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
@@ -17,6 +17,8 @@
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.handler.codec.mqtt.MqttMessage;
@@ -24,6 +26,8 @@ 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;
+import org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository;
+import org.apache.shenyu.protocol.mqtt.utils.MqttPacketIdGenerator;
/**
* mqtt transport handler.
@@ -48,7 +52,10 @@ public class MqttTransportHandler extends
ChannelInboundHandlerAdapter implement
@Override
public void operationComplete(final Future<? super Void> future) throws
Exception {
-
+ Channel channel = ((ChannelFuture) future).channel();
+ Singleton.INST.get(ChannelRepository.class).remove(channel);
+ Singleton.INST.get(SubscribeRepository.class).remove(channel);
+ MqttPacketIdGenerator.remove(channel);
}
}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PubRel.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PubRel.java
new file mode 100644
index 0000000000..516ceb96db
--- /dev/null
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PubRel.java
@@ -0,0 +1,41 @@
+/*
+ * 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.handler.codec.mqtt.MqttFixedHeader;
+import io.netty.handler.codec.mqtt.MqttMessage;
+import io.netty.handler.codec.mqtt.MqttMessageIdVariableHeader;
+import io.netty.handler.codec.mqtt.MqttQoS;
+
+import static io.netty.handler.codec.mqtt.MqttMessageType.PUBCOMP;
+
+/**
+ * The PUBREL message is the third message of the QoS 2 protocol flow,
+ * the server responds with PUBCOMP to release the packet id.
+ */
+public class PubRel extends MessageType {
+
+ @Override
+ public void pubRel(final ChannelHandlerContext ctx, final MqttMessage msg)
{
+ MqttMessageIdVariableHeader variableHeader =
(MqttMessageIdVariableHeader) msg.variableHeader();
+ MqttFixedHeader mqttFixedHeader = new MqttFixedHeader(PUBCOMP, false,
MqttQoS.AT_MOST_ONCE, false, 0);
+ MqttMessage mqttPubCompMessage = new MqttMessage(mqttFixedHeader,
variableHeader);
+ ctx.writeAndFlush(mqttPubCompMessage);
+ }
+}
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 c469b3cf22..db41e236f4 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
@@ -22,6 +22,7 @@ import io.netty.buffer.Unpooled;
import io.netty.channel.Channel;
import io.netty.channel.ChannelHandlerContext;
import io.netty.handler.codec.mqtt.MqttFixedHeader;
+import io.netty.handler.codec.mqtt.MqttMessage;
import io.netty.handler.codec.mqtt.MqttMessageIdVariableHeader;
import io.netty.handler.codec.mqtt.MqttPublishMessage;
import io.netty.handler.codec.mqtt.MqttQoS;
@@ -31,12 +32,14 @@ import
io.netty.handler.codec.mqtt.MqttPublishVariableHeader;
import org.apache.shenyu.common.utils.Singleton;
import org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository;
import org.apache.shenyu.protocol.mqtt.repositories.TopicRepository;
+import org.apache.shenyu.protocol.mqtt.utils.MqttPacketIdGenerator;
-import java.util.List;
+import java.util.Map;
import java.util.concurrent.CompletableFuture;
import static io.netty.channel.ChannelFutureListener.FIRE_EXCEPTION_ON_FAILURE;
import static io.netty.handler.codec.mqtt.MqttMessageType.PUBACK;
+import static io.netty.handler.codec.mqtt.MqttMessageType.PUBREC;
/**
* Publish message.
@@ -52,7 +55,6 @@ public class Publish extends MessageType {
String topic = msg.variableHeader().topicName();
ByteBuf payload = msg.payload();
String message = byteBufToString(payload);
- //// todo qos
MqttQoS mqttQoS = msg.fixedHeader().qosLevel();
if (msg.fixedHeader().isRetain()) {
if (payload.isReadable()) {
@@ -62,7 +64,7 @@ public class Publish extends MessageType {
}
}
int packetId = msg.variableHeader().packetId();
- CompletableFuture.runAsync(() -> send(topic, payload, packetId));
+ CompletableFuture.runAsync(() -> send(topic, payload, mqttQoS));
switch (mqttQoS.value()) {
case 0:
@@ -82,17 +84,10 @@ public class Publish extends MessageType {
}
/**
- * todo qos0.
- */
- private void qos0() {
-
- }
-
- /**
- * todo qos1.
+ * send PUBACK to the publisher for a qos1 publish.
*/
private void qos1(final ChannelHandlerContext ctx, final int packetId) {
- MqttFixedHeader mqttFixedHeader = new MqttFixedHeader(PUBACK, false,
MqttQoS.AT_LEAST_ONCE, false, 0);
+ MqttFixedHeader mqttFixedHeader = new MqttFixedHeader(PUBACK, false,
MqttQoS.AT_MOST_ONCE, false, 0);
MqttMessageIdVariableHeader mqttMsgIdVariableHeader =
MqttMessageIdVariableHeader.from(packetId);
MqttPubAckMessage mqttPubAckMessage = new
MqttPubAckMessage(mqttFixedHeader, mqttMsgIdVariableHeader);
@@ -100,14 +95,14 @@ public class Publish extends MessageType {
}
/**
- * todo qos2.
+ * send PUBREC to the publisher for a qos2 publish.
*/
private void qos2(final ChannelHandlerContext ctx, final int packetId) {
- MqttFixedHeader mqttFixedHeader = new MqttFixedHeader(PUBACK, false,
MqttQoS.EXACTLY_ONCE, false, 0);
+ MqttFixedHeader mqttFixedHeader = new MqttFixedHeader(PUBREC, false,
MqttQoS.AT_MOST_ONCE, false, 0);
MqttMessageIdVariableHeader mqttMsgIdVariableHeader =
MqttMessageIdVariableHeader.from(packetId);
- MqttPubAckMessage mqttPubAckMessage = new
MqttPubAckMessage(mqttFixedHeader, mqttMsgIdVariableHeader);
- ctx.writeAndFlush(mqttPubAckMessage);
+ MqttMessage mqttPubRecMessage = new MqttMessage(mqttFixedHeader,
mqttMsgIdVariableHeader);
+ ctx.writeAndFlush(mqttPubRecMessage);
}
private String byteBufToString(final ByteBuf byteBuf) {
@@ -120,16 +115,23 @@ public class Publish extends MessageType {
}
}
- private void send(final String topic, final ByteBuf payload, final int
packetId) {
- List<Channel> channels =
Singleton.INST.get(SubscribeRepository.class).get(topic);
+ private void send(final String topic, final ByteBuf payload, final MqttQoS
publishQoS) {
+ Map<Channel, MqttQoS> subscribers =
Singleton.INST.get(SubscribeRepository.class).get(topic);
//// todo thread pool
- channels.parallelStream().forEach(channel -> {
+ subscribers.entrySet().parallelStream().forEach(entry -> {
+ Channel channel = entry.getKey();
if (channel.isActive()) {
- MqttFixedHeader mqttFixedHeader = new
MqttFixedHeader(MqttMessageType.PUBLISH, false, MqttQoS.AT_MOST_ONCE, false, 0);
+ MqttQoS qos = minQoS(publishQoS, entry.getValue());
+ int packetId = MqttQoS.AT_MOST_ONCE == qos ? 0 :
MqttPacketIdGenerator.next(channel);
+ MqttFixedHeader mqttFixedHeader = new
MqttFixedHeader(MqttMessageType.PUBLISH, false, qos, false, 0);
MqttPublishVariableHeader mqttPublishVariableHeader = new
MqttPublishVariableHeader(topic, packetId);
- MqttPublishMessage mqttPublishMessage = new
MqttPublishMessage(mqttFixedHeader, mqttPublishVariableHeader,
Unpooled.wrappedBuffer(payload));
+ MqttPublishMessage mqttPublishMessage = new
MqttPublishMessage(mqttFixedHeader, mqttPublishVariableHeader,
Unpooled.wrappedBuffer(payload.retain()));
channel.writeAndFlush(mqttPublishMessage);
}
});
}
+
+ private static MqttQoS minQoS(final MqttQoS publishQoS, final MqttQoS
grantedQoS) {
+ return publishQoS.value() <= grantedQoS.value() ? publishQoS :
grantedQoS;
+ }
}
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 a0e21742f3..befd4a5c8e 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
@@ -18,35 +18,33 @@
package org.apache.shenyu.protocol.mqtt.repositories;
import io.netty.channel.Channel;
+import io.netty.handler.codec.mqtt.MqttQoS;
import io.netty.handler.codec.mqtt.MqttTopicSubscription;
-import org.apache.commons.collections4.CollectionUtils;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
+import java.util.Collections;
import java.util.List;
import java.util.Map;
-import java.util.Set;
+import java.util.Objects;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ConcurrentHashMap;
-import java.util.concurrent.CopyOnWriteArrayList;
-import java.util.concurrent.CopyOnWriteArraySet;
/**
* Topic and channel association.
*/
-public class SubscribeRepository implements BaseRepository<List<String>,
List<Channel>> {
+public class SubscribeRepository implements BaseRepository<List<String>,
Map<Channel, MqttQoS>> {
private static final Logger LOG =
LoggerFactory.getLogger(SubscribeRepository.class);
- private static final Map<String, List<Channel>> TOPIC_CHANNEL_FACTORY =
new ConcurrentHashMap<>();
+ private static final Map<String, Map<Channel, MqttQoS>>
TOPIC_CHANNEL_FACTORY = new ConcurrentHashMap<>();
@Override
- public void add(final List<String> topics, final List<Channel> channels) {
- CompletableFuture.runAsync(() -> topics.parallelStream().forEach(s -> {
- List<Channel> list = get(s);
- list.addAll(channels);
- TOPIC_CHANNEL_FACTORY.put(s, list);
- }));
+ public void add(final List<String> topics, final Map<Channel, MqttQoS>
channelQos) {
+ CompletableFuture.runAsync(() -> topics.parallelStream().forEach(topic
->
+ channelQos.forEach((channel, qos) -> TOPIC_CHANNEL_FACTORY
+ .computeIfAbsent(topic, key -> new
ConcurrentHashMap<>())
+ .merge(channel, qos, SubscribeRepository::maxQoS))));
}
/**
@@ -55,11 +53,11 @@ public class SubscribeRepository implements
BaseRepository<List<String>, List<Ch
* @param mqttTopicSubscription mqtt subscription info
*/
public void add(final Channel channel, final List<MqttTopicSubscription>
mqttTopicSubscription) {
- CompletableFuture.runAsync(() ->
mqttTopicSubscription.parallelStream().forEach(s -> {
- List<Channel> channels = get(s.topicName());
- channels.add(channel);
- TOPIC_CHANNEL_FACTORY.put(s.topicName(), channels);
- }));
+ CompletableFuture.runAsync(() -> mqttTopicSubscription.parallelStream()
+ .filter(s -> s.qualityOfService() != MqttQoS.FAILURE)
+ .forEach(s -> TOPIC_CHANNEL_FACTORY
+ .computeIfAbsent(s.topicName(), key -> new
ConcurrentHashMap<>())
+ .merge(channel, s.qualityOfService(),
SubscribeRepository::maxQoS)));
}
@Override
@@ -74,27 +72,41 @@ public class SubscribeRepository implements
BaseRepository<List<String>, List<Ch
*/
public void remove(final List<String> topics, final Channel channel) {
CompletableFuture.runAsync(() -> topics.parallelStream().forEach(topic
-> {
- List<Channel> channels = TOPIC_CHANNEL_FACTORY.get(topic);
- if (CollectionUtils.isNotEmpty(channels)) {
- channels.remove(channel);
+ Map<Channel, MqttQoS> subscribers =
TOPIC_CHANNEL_FACTORY.get(topic);
+ if (Objects.nonNull(subscribers)) {
+ subscribers.remove(channel);
}
}));
}
+ /**
+ * remove the channel from all topics it subscribed.
+ * @param channel channel
+ */
+ public void remove(final Channel channel) {
+ CompletableFuture.runAsync(() ->
TOPIC_CHANNEL_FACTORY.values().parallelStream()
+ .forEach(subscribers -> subscribers.remove(channel)));
+ }
+
@Override
- public List<Channel> get(final List<String> topics) {
- Set<Channel> channels = new CopyOnWriteArraySet<>();
- topics.parallelStream().forEach(s ->
channels.addAll(TOPIC_CHANNEL_FACTORY.get(s)));
- return new CopyOnWriteArrayList<>(channels);
+ public Map<Channel, MqttQoS> get(final List<String> topics) {
+ Map<Channel, MqttQoS> subscribers = new ConcurrentHashMap<>();
+ topics.parallelStream().forEach(topic ->
TOPIC_CHANNEL_FACTORY.getOrDefault(topic, Collections.emptyMap())
+ .forEach((channel, qos) -> subscribers.merge(channel, qos,
SubscribeRepository::maxQoS)));
+ return subscribers;
}
/**
- * get Channels.
+ * get subscriber channels with their granted qos.
* @param topic topic
- * @return Channels
+ * @return map of channel to granted qos
*/
- public List<Channel> get(final String topic) {
- return TOPIC_CHANNEL_FACTORY.getOrDefault(topic, new
CopyOnWriteArrayList<>());
+ public Map<Channel, MqttQoS> get(final String topic) {
+ return TOPIC_CHANNEL_FACTORY.getOrDefault(topic,
Collections.emptyMap());
+ }
+
+ private static MqttQoS maxQoS(final MqttQoS qos1, final MqttQoS qos2) {
+ return qos1.value() >= qos2.value() ? qos1 : qos2;
}
}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/utils/MqttPacketIdGenerator.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/utils/MqttPacketIdGenerator.java
new file mode 100644
index 0000000000..7dd1605451
--- /dev/null
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/utils/MqttPacketIdGenerator.java
@@ -0,0 +1,60 @@
+/*
+ * 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.utils;
+
+import io.netty.channel.Channel;
+
+import java.util.Collections;
+import java.util.Map;
+import java.util.WeakHashMap;
+import java.util.concurrent.atomic.AtomicInteger;
+
+/**
+ * Allocates packet identifiers for outbound messages from each channel's own
id space.
+ */
+public final class MqttPacketIdGenerator {
+
+ private static final int MIN_PACKET_ID = 1;
+
+ private static final int MAX_PACKET_ID = 0xFFFF;
+
+ // weak keys so channels closed without a DISCONNECT do not leak their id
space
+ private static final Map<Channel, AtomicInteger> CHANNEL_PACKET_ID_FACTORY
= Collections.synchronizedMap(new WeakHashMap<>());
+
+ private MqttPacketIdGenerator() {
+ }
+
+ /**
+ * get next packet id of the channel.
+ * @param channel channel
+ * @return next packet id
+ */
+ public static int next(final Channel channel) {
+ AtomicInteger packetId =
CHANNEL_PACKET_ID_FACTORY.computeIfAbsent(channel, key -> new AtomicInteger());
+ return packetId.updateAndGet(current -> current >= MAX_PACKET_ID ?
MIN_PACKET_ID : current + 1);
+ }
+
+ /**
+ * remove the channel packet id.
+ * @param channel channel
+ */
+ public static void remove(final Channel channel) {
+ CHANNEL_PACKET_ID_FACTORY.remove(channel);
+ }
+
+}
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 2f267af7e1..2b69c0c749 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
@@ -24,42 +24,80 @@ 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.MqttQoS;
+import io.netty.handler.codec.mqtt.MqttTopicSubscription;
import io.netty.handler.codec.mqtt.MqttVersion;
import org.apache.shenyu.common.utils.Singleton;
import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
-import org.junit.jupiter.api.AfterAll;
-import org.junit.jupiter.api.BeforeAll;
+import org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository;
+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 java.nio.charset.StandardCharsets;
+import java.time.Duration;
+import java.util.Collections;
+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;
/**
* Test cases for {@link MqttTransportHandler}.
*/
public final class MqttTransportHandlerTest {
+ private static final String TOPIC = "test/topic";
+
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 Duration TIMEOUT = Duration.ofSeconds(5);
+
+ private static final Duration POLL_INTERVAL = Duration.ofMillis(10);
+
+ /**
+ * The repositories keep their state in static maps shared with the other
test classes of this module,
+ * so they are registered before every test and released again afterwards.
+ */
+ private static final ChannelRepository CHANNEL_REPOSITORY = new
ChannelRepository();
- @BeforeAll
- static void setUp() {
- channelRepository = new ChannelRepository();
- Singleton.INST.single(ChannelRepository.class, channelRepository);
+ private static final SubscribeRepository SUBSCRIBE_REPOSITORY = new
SubscribeRepository();
+
+ private EmbeddedChannel registeredChannel;
+
+ @BeforeEach
+ public void setUp() {
+ Singleton.INST.single(ChannelRepository.class, CHANNEL_REPOSITORY);
+ Singleton.INST.single(SubscribeRepository.class, SUBSCRIBE_REPOSITORY);
new MqttContext().setUserName(USER_NAME);
new MqttContext().setPassword(PASSWORD);
+
+ registeredChannel = new EmbeddedChannel();
+ CHANNEL_REPOSITORY.add(registeredChannel, CLIENT_ID);
+ SUBSCRIBE_REPOSITORY.add(registeredChannel,
+ Collections.singletonList(new MqttTopicSubscription(TOPIC,
MqttQoS.AT_LEAST_ONCE)));
+ awaitAssert(() -> {
+ assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(registeredChannel));
+
assertTrue(SUBSCRIBE_REPOSITORY.get(TOPIC).containsKey(registeredChannel));
+ });
}
- @AfterAll
- static void tearDown() {
+ @AfterEach
+ public void tearDown() {
+ MqttPacketIdGenerator.remove(registeredChannel);
+ CHANNEL_REPOSITORY.remove(registeredChannel);
+ SUBSCRIBE_REPOSITORY.remove(registeredChannel);
+ awaitAssert(() ->
assertFalse(SUBSCRIBE_REPOSITORY.get(TOPIC).containsKey(registeredChannel)));
+
+ registeredChannel.finishAndReleaseAll();
+
new MqttContext().setUserName(null);
new MqttContext().setPassword(null);
}
@@ -69,13 +107,13 @@ public final class MqttTransportHandlerTest {
EmbeddedChannel channel = new EmbeddedChannel(new
MqttTransportHandler());
channel.writeInbound(connectMessage());
- assertEquals(CLIENT_ID, channelRepository.get(channel));
+ assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(channel));
channel.writeInbound(connectMessage());
channel.runPendingTasks();
assertFalse(channel.isActive());
- assertNull(channelRepository.get(channel));
+ assertNull(CHANNEL_REPOSITORY.get(channel));
channel.finishAndReleaseAll();
}
@@ -84,13 +122,27 @@ public final class MqttTransportHandlerTest {
EmbeddedChannel channel = new EmbeddedChannel(new
MqttTransportHandler());
channel.writeInbound(connectMessage());
- assertEquals(CLIENT_ID, channelRepository.get(channel));
+ assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(channel));
channel.close();
channel.runPendingTasks();
assertFalse(channel.isActive());
- assertNull(channelRepository.get(channel));
+ assertNull(CHANNEL_REPOSITORY.get(channel));
+ channel.finishAndReleaseAll();
+ }
+
+ @Test
+ public void nonMqttMessageClosesConnectedChannel() {
+ EmbeddedChannel channel = new EmbeddedChannel(new
MqttTransportHandler());
+
+ channel.writeInbound(connectMessage());
+ assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(channel));
+
+ channel.writeInbound("not-a-mqtt-message");
+
+ assertFalse(channel.isActive());
+ assertNull(CHANNEL_REPOSITORY.get(channel));
channel.finishAndReleaseAll();
}
@@ -104,4 +156,24 @@ public final class MqttTransportHandlerTest {
return new MqttConnectMessage(fixedHeader, variableHeader, payload);
}
+ @Test
+ public void testOperationCompleteCleansRepositoriesOnClose() throws
Exception {
+ assertEquals(1, MqttPacketIdGenerator.next(registeredChannel));
+
+ new
MqttTransportHandler().operationComplete(registeredChannel.closeFuture());
+
+ awaitAssert(() ->
assertNull(CHANNEL_REPOSITORY.get(registeredChannel)));
+ awaitAssert(() ->
assertFalse(SUBSCRIBE_REPOSITORY.get(TOPIC).containsKey(registeredChannel)));
+ assertEquals(1, MqttPacketIdGenerator.next(registeredChannel));
+ }
+
+ /**
+ * The repositories mutate their state asynchronously on the common pool,
+ * so assertions are retried until the mutation becomes visible.
+ *
+ * @param assertion assertion to retry
+ */
+ private void awaitAssert(final ThrowingRunnable assertion) {
+
await().atMost(TIMEOUT).pollInterval(POLL_INTERVAL).untilAsserted(assertion);
+ }
}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PubRelTest.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PubRelTest.java
new file mode 100644
index 0000000000..1acd34c6e4
--- /dev/null
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PubRelTest.java
@@ -0,0 +1,55 @@
+/*
+ * 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.handler.codec.mqtt.MqttFixedHeader;
+import io.netty.handler.codec.mqtt.MqttMessage;
+import io.netty.handler.codec.mqtt.MqttMessageIdVariableHeader;
+import io.netty.handler.codec.mqtt.MqttMessageType;
+import io.netty.handler.codec.mqtt.MqttQoS;
+import org.junit.jupiter.api.Test;
+import org.mockito.ArgumentCaptor;
+
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.verify;
+
+/**
+ * Test cases for PubRel.
+ */
+public class PubRelTest {
+
+ private static final int PACKET_ID = 12345;
+
+ private final ChannelHandlerContext ctx =
mock(ChannelHandlerContext.class);
+
+ @Test
+ public void testPubRelRespondsWithPubComp() {
+ MqttFixedHeader pubRelFixedHeader = new
MqttFixedHeader(MqttMessageType.PUBREL, false, MqttQoS.AT_LEAST_ONCE, false, 0);
+ MqttMessage pubRel = new MqttMessage(pubRelFixedHeader,
MqttMessageIdVariableHeader.from(PACKET_ID));
+ new PubRel().pubRel(ctx, pubRel);
+
+ ArgumentCaptor<MqttMessage> captor =
ArgumentCaptor.forClass(MqttMessage.class);
+ verify(ctx).writeAndFlush(captor.capture());
+ MqttMessage pubComp = captor.getValue();
+ assertEquals(MqttMessageType.PUBCOMP,
pubComp.fixedHeader().messageType());
+ assertEquals(MqttQoS.AT_MOST_ONCE, pubComp.fixedHeader().qosLevel());
+ assertEquals(PACKET_ID, ((MqttMessageIdVariableHeader)
pubComp.variableHeader()).messageId());
+ }
+}
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 9b5d56ef01..487208d800 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
@@ -17,7 +17,9 @@
package org.apache.shenyu.protocol.mqtt;
+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;
@@ -25,26 +27,45 @@ 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.MqttMessage;
+import io.netty.handler.codec.mqtt.MqttMessageIdVariableHeader;
import io.netty.handler.codec.mqtt.MqttMessageType;
+import io.netty.handler.codec.mqtt.MqttPubAckMessage;
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.MqttTopicSubscription;
import io.netty.handler.codec.mqtt.MqttVersion;
import io.netty.util.CharsetUtil;
+import io.netty.util.ReferenceCountUtil;
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.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.mockito.ArgumentCaptor;
import java.nio.charset.StandardCharsets;
import java.time.Duration;
+import java.util.Arrays;
+import java.util.Collections;
+import java.util.List;
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.any;
+import static org.mockito.Mockito.never;
+import static org.mockito.Mockito.timeout;
+import static org.mockito.Mockito.verify;
+import static org.mockito.Mockito.when;
/**
* Test cases for {@link Publish}.
@@ -61,81 +82,286 @@ public final class PublishTest {
private static final String END_TO_END_TOPIC = "test/end-to-end";
+ private static final String TOPIC = "test/topic";
+
+ private static final List<String> ALL_TOPICS = Arrays.asList(
+ RETAINED_TOPIC, NON_RETAINED_TOPIC, CLEARED_TOPIC,
UNCONNECTED_TOPIC, END_TO_END_TOPIC, TOPIC);
+
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;
+ private static final String PAYLOAD = "hello";
+
+ private static final int PUBLISHER_PACKET_ID = 12345;
+
+ private static final Duration TIMEOUT = Duration.ofSeconds(5);
+
+ private static final Duration POLL_INTERVAL = Duration.ofMillis(10);
+
+ /**
+ * The repositories hold their state in static maps shared with the other
test classes of this module,
+ * so they are released around every test.
+ */
+ private static final SubscribeRepository SUBSCRIBE_REPOSITORY = new
SubscribeRepository();
+
+ private static final TopicRepository TOPIC_REPOSITORY = new
TopicRepository();
+
+ private static final ChannelRepository CHANNEL_REPOSITORY = new
ChannelRepository();
+
+ private final Channel subscriberChannel = mock(Channel.class);
+
+ private final Channel otherSubscriberChannel = mock(Channel.class);
+
+ private EmbeddedChannel publisherChannel;
- @BeforeAll
- static void setUp() {
- topicRepository = new TopicRepository();
- Singleton.INST.single(TopicRepository.class, topicRepository);
- Singleton.INST.single(SubscribeRepository.class, new
SubscribeRepository());
- Singleton.INST.single(ChannelRepository.class, new
ChannelRepository());
+ @BeforeEach
+ public void setUp() {
+ Singleton.INST.single(SubscribeRepository.class, SUBSCRIBE_REPOSITORY);
+ Singleton.INST.single(TopicRepository.class, TOPIC_REPOSITORY);
+ Singleton.INST.single(ChannelRepository.class, CHANNEL_REPOSITORY);
new MqttContext().setUserName(USER_NAME);
new MqttContext().setPassword(PASSWORD);
+
+ when(subscriberChannel.isActive()).thenReturn(true);
+ when(otherSubscriberChannel.isActive()).thenReturn(true);
+
+ clearSharedState();
+ publisherChannel = channel(true);
}
- @AfterAll
- static void tearDown() {
+ @AfterEach
+ public void tearDown() {
+ MqttPacketIdGenerator.remove(subscriberChannel);
+ MqttPacketIdGenerator.remove(otherSubscriberChannel);
+ publisherChannel.finishAndReleaseAll();
+ clearSharedState();
+
new MqttContext().setUserName(null);
new MqttContext().setPassword(null);
}
@Test
public void retainedPublishStoresMessage() {
- new Publish().publish(connectedContext(),
publishMessage(RETAINED_TOPIC, "hello", true));
- await().atMost(Duration.ofSeconds(5))
- .until(() ->
"hello".equals(topicRepository.get(RETAINED_TOPIC)));
+ new Publish().publish(publisherContext(),
publishMessage(RETAINED_TOPIC, PAYLOAD, true));
+ awaitAssert(() -> assertEquals(PAYLOAD,
TOPIC_REPOSITORY.get(RETAINED_TOPIC)));
}
@Test
public void nonRetainedPublishDoesNotStoreMessage() {
- new Publish().publish(connectedContext(),
publishMessage(NON_RETAINED_TOPIC, "hello", false));
- assertNull(topicRepository.get(NON_RETAINED_TOPIC));
+ new Publish().publish(publisherContext(),
publishMessage(NON_RETAINED_TOPIC, PAYLOAD, false));
+ assertNull(TOPIC_REPOSITORY.get(NON_RETAINED_TOPIC));
}
@Test
- public void publishBeforeConnectClosesChannel() {
- EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
- ChannelHandlerContext ctx = channel.pipeline().lastContext();
+ public void retainedPublishReadsMessageFromDirectPayload() {
+ ByteBuf payload =
Unpooled.directBuffer().writeBytes(PAYLOAD.getBytes(StandardCharsets.UTF_8));
+ try {
+ new Publish().publish(publisherContext(),
publishMessage(RETAINED_TOPIC, payload, true));
+ awaitAssert(() -> assertEquals(PAYLOAD,
TOPIC_REPOSITORY.get(RETAINED_TOPIC)));
+ } finally {
+ payload.release();
+ }
+ }
- new Publish().publish(ctx, publishMessage(UNCONNECTED_TOPIC, "hello",
true));
+ @Test
+ public void publishBeforeConnectClosesChannel() {
+ EmbeddedChannel channel = channel(false);
+ new Publish().publish(channel.pipeline().lastContext(),
publishMessage(UNCONNECTED_TOPIC, PAYLOAD, true));
channel.runPendingTasks();
+
assertFalse(channel.isActive());
- assertNull(topicRepository.get(UNCONNECTED_TOPIC));
+ assertNull(TOPIC_REPOSITORY.get(UNCONNECTED_TOPIC));
+ channel.finishAndReleaseAll();
}
@Test
public void publishAfterConnectOnSameChannelIsAccepted() {
- EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
+ EmbeddedChannel channel = channel(false);
ChannelHandlerContext ctx = channel.pipeline().lastContext();
new Connect().connect(ctx, connectMessage());
- new Publish().publish(ctx, publishMessage(END_TO_END_TOPIC, "hello",
true));
+ new Publish().publish(ctx, publishMessage(END_TO_END_TOPIC, PAYLOAD,
true));
+
+ awaitAssert(() -> assertEquals(PAYLOAD,
TOPIC_REPOSITORY.get(END_TO_END_TOPIC)));
+ assertEquals(CLIENT_ID, CHANNEL_REPOSITORY.get(channel));
- await().atMost(Duration.ofSeconds(5))
- .until(() ->
"hello".equals(topicRepository.get(END_TO_END_TOPIC)));
+ CHANNEL_REPOSITORY.remove(channel);
+ channel.finishAndReleaseAll();
}
@Test
public void zeroByteRetainedPublishClearsRetainedMessage() {
Publish publish = new Publish();
- publish.publish(connectedContext(), publishMessage(CLEARED_TOPIC,
"hello", true));
- await().atMost(Duration.ofSeconds(5))
- .until(() ->
"hello".equals(topicRepository.get(CLEARED_TOPIC)));
- publish.publish(connectedContext(), publishMessage(CLEARED_TOPIC, "",
true));
- assertNull(topicRepository.get(CLEARED_TOPIC));
+ ChannelHandlerContext ctx = publisherContext();
+
+ publish.publish(ctx, publishMessage(CLEARED_TOPIC, PAYLOAD, true));
+ awaitAssert(() -> assertEquals(PAYLOAD,
TOPIC_REPOSITORY.get(CLEARED_TOPIC)));
+
+ publish.publish(ctx, publishMessage(CLEARED_TOPIC, "", true));
+ assertNull(TOPIC_REPOSITORY.get(CLEARED_TOPIC));
+ }
+
+ @Test
+ public void testPublishDeliversAtGrantedQosWithOwnPacketId() {
+ addSubscriber(subscriberChannel, MqttQoS.AT_LEAST_ONCE);
+ publishToSubscribers(MqttQoS.EXACTLY_ONCE);
+
+ MqttPublishMessage message = captureMessage(subscriberChannel);
+ assertEquals(MqttQoS.AT_LEAST_ONCE, message.fixedHeader().qosLevel());
+ assertEquals(1, message.variableHeader().packetId());
+ assertEquals(TOPIC, message.variableHeader().topicName());
+ assertEquals(PAYLOAD, message.payload().toString(CharsetUtil.UTF_8));
+ }
+
+ @Test
+ public void testPublishDeliversQos0SubscriberWithZeroPacketId() {
+ addSubscriber(subscriberChannel, MqttQoS.AT_MOST_ONCE);
+ publishToSubscribers(MqttQoS.EXACTLY_ONCE);
+
+ MqttPublishMessage message = captureMessage(subscriberChannel);
+ assertEquals(MqttQoS.AT_MOST_ONCE, message.fixedHeader().qosLevel());
+ assertEquals(0, message.variableHeader().packetId());
+ }
+
+ @Test
+ public void testPublishQos0FanOutDeliversAtMostOnce() {
+ addSubscriber(subscriberChannel, MqttQoS.EXACTLY_ONCE);
+ publishToSubscribers(MqttQoS.AT_MOST_ONCE);
+
+ MqttPublishMessage message = captureMessage(subscriberChannel);
+ assertEquals(MqttQoS.AT_MOST_ONCE, message.fixedHeader().qosLevel());
+ assertEquals(0, message.variableHeader().packetId());
+ }
+
+ @Test
+ public void testPublishSkipsInactiveSubscriber() {
+ addSubscriber(subscriberChannel, MqttQoS.AT_LEAST_ONCE);
+ addSubscriber(otherSubscriberChannel, MqttQoS.AT_LEAST_ONCE);
+ when(subscriberChannel.isActive()).thenReturn(false);
+
+ publishToSubscribers(MqttQoS.AT_LEAST_ONCE);
+
+ assertEquals(MqttQoS.AT_LEAST_ONCE,
captureMessage(otherSubscriberChannel).fixedHeader().qosLevel());
+ verify(subscriberChannel,
never()).writeAndFlush(any(MqttPublishMessage.class));
}
- private ChannelHandlerContext connectedContext() {
+ @Test
+ public void testPublishAllocatesPacketIdFromSubscriberIdSpace() {
+ addSubscriber(subscriberChannel, MqttQoS.EXACTLY_ONCE);
+ addSubscriber(otherSubscriberChannel, MqttQoS.EXACTLY_ONCE);
+
+ publishToSubscribers(MqttQoS.EXACTLY_ONCE);
+ publishToSubscribers(MqttQoS.EXACTLY_ONCE);
+
+ List<MqttPublishMessage> messages = captureMessages(subscriberChannel,
2);
+ assertEquals(1, messages.get(0).variableHeader().packetId());
+ assertEquals(2, messages.get(1).variableHeader().packetId());
+
+ List<MqttPublishMessage> otherMessages =
captureMessages(otherSubscriberChannel, 2);
+ assertEquals(1, otherMessages.get(0).variableHeader().packetId());
+ assertEquals(2, otherMessages.get(1).variableHeader().packetId());
+ }
+
+ @Test
+ public void testPublishFanOutRetainsPayloadPerSubscriber() {
+ addSubscriber(subscriberChannel, MqttQoS.AT_LEAST_ONCE);
+ addSubscriber(otherSubscriberChannel, MqttQoS.AT_LEAST_ONCE);
+ ByteBuf payload = Unpooled.copiedBuffer(PAYLOAD, CharsetUtil.UTF_8);
+ try {
+ publishToSubscribers(MqttQoS.AT_LEAST_ONCE, payload);
+ awaitAssert(() -> assertEquals(3, payload.refCnt()));
+
+ MqttPublishMessage delivered = captureMessage(subscriberChannel);
+ MqttPublishMessage otherDelivered =
captureMessage(otherSubscriberChannel);
+ assertEquals(PAYLOAD,
delivered.payload().toString(CharsetUtil.UTF_8));
+ assertEquals(PAYLOAD,
otherDelivered.payload().toString(CharsetUtil.UTF_8));
+
+ ReferenceCountUtil.release(delivered);
+ ReferenceCountUtil.release(otherDelivered);
+ assertEquals(1, payload.refCnt());
+ } finally {
+ ReferenceCountUtil.release(payload);
+ }
+ }
+
+ @Test
+ public void testPublishQos1SendsPubAckToPublisher() {
+ publishToSubscribers(MqttQoS.AT_LEAST_ONCE);
+
+ MqttPubAckMessage pubAck = awaitOutbound(publisherChannel);
+ assertEquals(PUBLISHER_PACKET_ID, pubAck.variableHeader().messageId());
+ assertEquals(MqttQoS.AT_MOST_ONCE, pubAck.fixedHeader().qosLevel());
+ }
+
+ @Test
+ public void testPublishQos2SendsPubRecToPublisher() {
+ publishToSubscribers(MqttQoS.EXACTLY_ONCE);
+
+ MqttMessage pubRec = awaitOutbound(publisherChannel);
+ assertEquals(MqttMessageType.PUBREC,
pubRec.fixedHeader().messageType());
+ assertEquals(MqttQoS.AT_MOST_ONCE, pubRec.fixedHeader().qosLevel());
+ assertEquals(PUBLISHER_PACKET_ID, ((MqttMessageIdVariableHeader)
pubRec.variableHeader()).messageId());
+ }
+
+ /**
+ * Creates a real channel, optionally already connected, so that {@link
Publish} can
+ * read the connection attribute and close the channel as it does in
production.
+ *
+ * @param connected whether the channel completed the CONNECT handshake
+ * @return the channel
+ */
+ private EmbeddedChannel channel(final boolean connected) {
EmbeddedChannel channel = new EmbeddedChannel(new
ChannelInboundHandlerAdapter());
- new MessageType().setConnected(channel, true);
- return channel.pipeline().lastContext();
+ if (connected) {
+ new MessageType().setConnected(channel, true);
+ }
+ return channel;
+ }
+
+ private ChannelHandlerContext publisherContext() {
+ return publisherChannel.pipeline().lastContext();
+ }
+
+ private void addSubscriber(final Channel channel, final MqttQoS qos) {
+ SUBSCRIBE_REPOSITORY.add(channel, Collections.singletonList(new
MqttTopicSubscription(TOPIC, qos)));
+ awaitAssert(() -> assertEquals(qos,
SUBSCRIBE_REPOSITORY.get(TOPIC).get(channel)));
+ }
+
+ private void publishToSubscribers(final MqttQoS qos) {
+ publishToSubscribers(qos, Unpooled.copiedBuffer(PAYLOAD,
CharsetUtil.UTF_8));
+ }
+
+ private void publishToSubscribers(final MqttQoS qos, final ByteBuf
payload) {
+ MqttFixedHeader fixedHeader = new
MqttFixedHeader(MqttMessageType.PUBLISH, false, qos, false, 0);
+ MqttPublishVariableHeader variableHeader = new
MqttPublishVariableHeader(TOPIC, PUBLISHER_PACKET_ID);
+ new Publish().publish(publisherContext(), new
MqttPublishMessage(fixedHeader, variableHeader, payload));
+ }
+
+ private MqttPublishMessage captureMessage(final Channel channel) {
+ return captureMessages(channel, 1).get(0);
+ }
+
+ private List<MqttPublishMessage> captureMessages(final Channel channel,
final int times) {
+ ArgumentCaptor<MqttPublishMessage> captor =
ArgumentCaptor.forClass(MqttPublishMessage.class);
+ verify(channel,
timeout(TIMEOUT.toMillis()).times(times)).writeAndFlush(captor.capture());
+ return captor.getAllValues();
+ }
+
+ /**
+ * Polls the messages written back to the publisher, such as PUBACK and
PUBREC.
+ *
+ * @param channel the publisher channel
+ * @param <T> the expected message type
+ * @return the first outbound message
+ */
+ private <T> T awaitOutbound(final EmbeddedChannel channel) {
+ channel.runPendingTasks();
+ awaitAssert(() -> assertFalse(channel.outboundMessages().isEmpty()));
+ return channel.readOutbound();
}
private MqttConnectMessage connectMessage() {
@@ -149,8 +375,28 @@ public final class PublishTest {
}
private MqttPublishMessage publishMessage(final String topic, final String
payload, final boolean retain) {
+ return publishMessage(topic, Unpooled.copiedBuffer(payload,
CharsetUtil.UTF_8), retain);
+ }
+
+ private MqttPublishMessage publishMessage(final String topic, final
ByteBuf payload, final boolean retain) {
MqttFixedHeader fixedHeader = new
MqttFixedHeader(MqttMessageType.PUBLISH, false, MqttQoS.AT_MOST_ONCE, retain,
0);
MqttPublishVariableHeader variableHeader = new
MqttPublishVariableHeader(topic, 1);
- return new MqttPublishMessage(fixedHeader, variableHeader,
Unpooled.copiedBuffer(payload, CharsetUtil.UTF_8));
+ return new MqttPublishMessage(fixedHeader, variableHeader, payload);
+ }
+
+ /**
+ * Subscriptions are registered asynchronously on the common pool,
+ * so assertions are retried until the mutation becomes visible.
+ *
+ * @param assertion assertion to retry
+ */
+ private void awaitAssert(final ThrowingRunnable assertion) {
+
await().atMost(TIMEOUT).pollInterval(POLL_INTERVAL).untilAsserted(assertion);
+ }
+
+ private void clearSharedState() {
+ ALL_TOPICS.forEach(TOPIC_REPOSITORY::remove);
+ SUBSCRIBE_REPOSITORY.remove(ALL_TOPICS);
+ awaitAssert(() -> ALL_TOPICS.forEach(topic ->
assertTrue(SUBSCRIBE_REPOSITORY.get(topic).isEmpty())));
}
}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepositoryTest.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepositoryTest.java
index 90936503d4..77e674f387 100644
---
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepositoryTest.java
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepositoryTest.java
@@ -20,15 +20,25 @@ package org.apache.shenyu.protocol.mqtt.repositories;
import io.netty.channel.Channel;
import io.netty.handler.codec.mqtt.MqttQoS;
import io.netty.handler.codec.mqtt.MqttTopicSubscription;
+import org.apache.shenyu.common.utils.Singleton;
+import org.awaitility.core.ThrowingRunnable;
+import org.junit.jupiter.api.AfterEach;
+import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import java.time.Duration;
+import java.util.Arrays;
import java.util.Collections;
+import java.util.List;
+import java.util.Map;
+import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ForkJoinPool;
import java.util.concurrent.TimeUnit;
import static org.awaitility.Awaitility.await;
import static org.junit.jupiter.api.Assertions.assertDoesNotThrow;
+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.Mockito.mock;
@@ -37,36 +47,180 @@ import static org.mockito.Mockito.mock;
*/
public final class SubscribeRepositoryTest {
- private static final String EXISTING_TOPIC = "test/existing-topic";
-
private static final String ABSENT_TOPIC = "test/absent-topic";
- private static final String KEPT_TOPIC = "test/kept-topic";
+ private static final String TOPIC = "test/topic";
+
+ private static final String OTHER_TOPIC = "test/other-topic";
+
+ private static final List<String> ALL_TOPICS = Arrays.asList(ABSENT_TOPIC,
TOPIC, OTHER_TOPIC);
+
+ private static final Duration TIMEOUT = Duration.ofSeconds(5);
+
+ private static final Duration POLL_INTERVAL = Duration.ofMillis(10);
+
+ private SubscribeRepository repository;
+
+ private Channel channel;
+
+ private Channel otherChannel;
+
+ @BeforeEach
+ public void setUp() {
+ repository = new SubscribeRepository();
+ channel = mock(Channel.class);
+ otherChannel = mock(Channel.class);
+ Singleton.INST.single(SubscribeRepository.class, repository);
+ clearAllTopics();
+ }
+
+ @AfterEach
+ public void tearDown() {
+ clearAllTopics();
+ }
@Test
- public void removeRemovesChannelFromExistingTopic() {
- SubscribeRepository repository = new SubscribeRepository();
- Channel channel = mock(Channel.class);
- repository.add(channel, Collections.singletonList(new
MqttTopicSubscription(EXISTING_TOPIC, MqttQoS.AT_MOST_ONCE)));
- await().atMost(Duration.ofSeconds(5)).until(() ->
repository.get(EXISTING_TOPIC).contains(channel));
+ public void testAddStoresGrantedQosPerTopic() {
+ repository.add(channel, Collections.singletonList(new
MqttTopicSubscription(TOPIC, MqttQoS.AT_LEAST_ONCE)));
+ awaitAssert(() -> assertEquals(MqttQoS.AT_LEAST_ONCE,
repository.get(TOPIC).get(channel)));
+ }
+
+ @Test
+ public void testAddRegistersEverySubscribedTopic() {
+ repository.add(channel, Arrays.asList(
+ new MqttTopicSubscription(TOPIC, MqttQoS.AT_MOST_ONCE),
+ new MqttTopicSubscription(OTHER_TOPIC, MqttQoS.EXACTLY_ONCE)));
+ awaitAssert(() -> {
+ assertEquals(MqttQoS.AT_MOST_ONCE,
repository.get(TOPIC).get(channel));
+ assertEquals(MqttQoS.EXACTLY_ONCE,
repository.get(OTHER_TOPIC).get(channel));
+ });
+ }
- repository.remove(Collections.singletonList(EXISTING_TOPIC), channel);
+ @Test
+ public void testAddKeepsMaxQosForOverlappingSubscription() {
+ repository.add(channel, Arrays.asList(
+ new MqttTopicSubscription(TOPIC, MqttQoS.AT_LEAST_ONCE),
+ new MqttTopicSubscription(TOPIC, MqttQoS.EXACTLY_ONCE)));
+ awaitAssert(() -> assertEquals(MqttQoS.EXACTLY_ONCE,
repository.get(TOPIC).get(channel)));
- await().atMost(Duration.ofSeconds(5)).until(() ->
repository.get(EXISTING_TOPIC).isEmpty());
+ repository.add(channel, Collections.singletonList(new
MqttTopicSubscription(TOPIC, MqttQoS.AT_MOST_ONCE)));
+ awaitRepositoryIdle();
+ assertEquals(MqttQoS.EXACTLY_ONCE, repository.get(TOPIC).get(channel));
}
@Test
- public void removeAbsentTopicDoesNotThrow() {
- SubscribeRepository repository = new SubscribeRepository();
- Channel channel = mock(Channel.class);
- repository.add(channel, Collections.singletonList(new
MqttTopicSubscription(KEPT_TOPIC, MqttQoS.AT_MOST_ONCE)));
- await().atMost(Duration.ofSeconds(5)).until(() ->
repository.get(KEPT_TOPIC).contains(channel));
+ public void testAddIgnoresFailureSubscription() {
+ repository.add(channel, Collections.singletonList(new
MqttTopicSubscription(TOPIC, MqttQoS.FAILURE)));
+ awaitRepositoryIdle();
+ assertTrue(repository.get(TOPIC).isEmpty());
+ assertTrue(repository.get(ALL_TOPICS).isEmpty());
+ }
+
+ @Test
+ public void testAddTopicsWithChannelQosMap() {
+ Map<Channel, MqttQoS> channelQos = new ConcurrentHashMap<>();
+ channelQos.put(channel, MqttQoS.AT_MOST_ONCE);
+ channelQos.put(otherChannel, MqttQoS.EXACTLY_ONCE);
+ repository.add(Arrays.asList(TOPIC, OTHER_TOPIC), channelQos);
+ awaitAssert(() -> {
+ assertEquals(MqttQoS.AT_MOST_ONCE,
repository.get(TOPIC).get(channel));
+ assertEquals(MqttQoS.EXACTLY_ONCE,
repository.get(TOPIC).get(otherChannel));
+ assertEquals(MqttQoS.AT_MOST_ONCE,
repository.get(OTHER_TOPIC).get(channel));
+ assertEquals(MqttQoS.EXACTLY_ONCE,
repository.get(OTHER_TOPIC).get(otherChannel));
+ });
+ }
+
+ @Test
+ public void testGetMergesSubscribersOfEveryTopic() {
+ repository.add(channel, Collections.singletonList(new
MqttTopicSubscription(TOPIC, MqttQoS.AT_MOST_ONCE)));
+ repository.add(channel, Collections.singletonList(new
MqttTopicSubscription(OTHER_TOPIC, MqttQoS.EXACTLY_ONCE)));
+ repository.add(otherChannel, Collections.singletonList(new
MqttTopicSubscription(OTHER_TOPIC, MqttQoS.AT_LEAST_ONCE)));
+ awaitAssert(() -> {
+ Map<Channel, MqttQoS> subscribers =
repository.get(Arrays.asList(TOPIC, OTHER_TOPIC));
+ assertEquals(2, subscribers.size());
+ assertEquals(MqttQoS.EXACTLY_ONCE, subscribers.get(channel));
+ assertEquals(MqttQoS.AT_LEAST_ONCE, subscribers.get(otherChannel));
+ });
+ }
+
+ @Test
+ public void testGetAbsentTopicReturnsNoSubscribers() {
+
assertTrue(repository.get(Collections.singletonList(ABSENT_TOPIC)).isEmpty());
+ }
+
+ @Test
+ public void testRemoveChannelFromTopic() {
+ repository.add(channel, Arrays.asList(
+ new MqttTopicSubscription(TOPIC, MqttQoS.AT_MOST_ONCE),
+ new MqttTopicSubscription(OTHER_TOPIC, MqttQoS.AT_MOST_ONCE)));
+ repository.add(otherChannel, Collections.singletonList(new
MqttTopicSubscription(TOPIC, MqttQoS.AT_LEAST_ONCE)));
+ awaitAssert(() -> {
+ assertEquals(MqttQoS.AT_MOST_ONCE,
repository.get(TOPIC).get(channel));
+ assertEquals(MqttQoS.AT_LEAST_ONCE,
repository.get(TOPIC).get(otherChannel));
+ });
+
+ repository.remove(Collections.singletonList(TOPIC), channel);
+
+ awaitAssert(() -> {
+ assertFalse(repository.get(TOPIC).containsKey(channel));
+ assertEquals(MqttQoS.AT_LEAST_ONCE,
repository.get(TOPIC).get(otherChannel));
+ assertEquals(MqttQoS.AT_MOST_ONCE,
repository.get(OTHER_TOPIC).get(channel));
+ });
+ }
+
+ @Test
+ public void testRemoveChannelFromEveryTopic() {
+ repository.add(channel, Arrays.asList(
+ new MqttTopicSubscription(TOPIC, MqttQoS.AT_MOST_ONCE),
+ new MqttTopicSubscription(OTHER_TOPIC, MqttQoS.AT_MOST_ONCE)));
+ repository.add(otherChannel, Collections.singletonList(new
MqttTopicSubscription(TOPIC, MqttQoS.AT_LEAST_ONCE)));
+ awaitAssert(() -> assertEquals(MqttQoS.AT_LEAST_ONCE,
repository.get(TOPIC).get(otherChannel)));
+
+ repository.remove(channel);
+
+ awaitAssert(() -> {
+ assertFalse(repository.get(TOPIC).containsKey(channel));
+ assertTrue(repository.get(OTHER_TOPIC).isEmpty());
+ assertEquals(MqttQoS.AT_LEAST_ONCE,
repository.get(TOPIC).get(otherChannel));
+ });
+ }
+
+ @Test
+ public void testRemoveAbsentTopicDoesNotThrow() {
+ repository.add(channel, Collections.singletonList(new
MqttTopicSubscription(TOPIC, MqttQoS.AT_MOST_ONCE)));
+ awaitAssert(() -> assertEquals(MqttQoS.AT_MOST_ONCE,
repository.get(TOPIC).get(channel)));
assertDoesNotThrow(() ->
repository.remove(Collections.singletonList(ABSENT_TOPIC), channel));
- await().atMost(Duration.ofSeconds(5))
- .until(() -> ForkJoinPool.commonPool().awaitQuiescence(1,
TimeUnit.SECONDS));
+ awaitRepositoryIdle();
assertTrue(repository.get(ABSENT_TOPIC).isEmpty());
- assertTrue(repository.get(KEPT_TOPIC).contains(channel));
+ assertEquals(MqttQoS.AT_MOST_ONCE, repository.get(TOPIC).get(channel));
+ }
+
+ /**
+ * The repository mutates its state asynchronously on the common pool,
+ * so assertions have to be retried until the mutation becomes visible.
+ *
+ * @param assertion assertion to retry
+ */
+ private void awaitAssert(final ThrowingRunnable assertion) {
+
await().atMost(TIMEOUT).pollInterval(POLL_INTERVAL).untilAsserted(assertion);
+ }
+
+ /**
+ * Waits until the repository finished all pending asynchronous mutations.
+ * Required to assert that a mutation did <em>not</em> change the shared
state.
+ */
+ private void awaitRepositoryIdle() {
+
assertTrue(ForkJoinPool.commonPool().awaitQuiescence(TIMEOUT.toMillis(),
TimeUnit.MILLISECONDS));
+ }
+
+ /**
+ * The repository keeps its state in a static map which is shared by every
instance and
+ * by the other test classes of this module, so the topics used here are
released around every test.
+ */
+ private void clearAllTopics() {
+ repository.remove(ALL_TOPICS);
+ awaitAssert(() -> ALL_TOPICS.forEach(topic ->
assertTrue(repository.get(topic).isEmpty())));
}
}
diff --git
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/utils/MqttPacketIdGeneratorTest.java
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/utils/MqttPacketIdGeneratorTest.java
new file mode 100644
index 0000000000..3360361ea2
--- /dev/null
+++
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/utils/MqttPacketIdGeneratorTest.java
@@ -0,0 +1,77 @@
+/*
+ * 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.utils;
+
+import io.netty.channel.Channel;
+import org.junit.jupiter.api.AfterEach;
+import org.junit.jupiter.api.Test;
+
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.mockito.Mockito.mock;
+
+/**
+ * Test cases for MqttPacketIdGenerator.
+ */
+public class MqttPacketIdGeneratorTest {
+
+ private static final int MAX_PACKET_ID = 0xFFFF;
+
+ private final Channel channel = mock(Channel.class);
+
+ @AfterEach
+ public void tearDown() {
+ MqttPacketIdGenerator.remove(channel);
+ }
+
+ @Test
+ public void testNextAllocatesSequentialPacketIds() {
+ assertEquals(1, MqttPacketIdGenerator.next(channel));
+ assertEquals(2, MqttPacketIdGenerator.next(channel));
+ assertEquals(3, MqttPacketIdGenerator.next(channel));
+ }
+
+ @Test
+ public void testNextKeepsIndependentIdSpacesPerChannel() {
+ Channel otherChannel = mock(Channel.class);
+ try {
+ assertEquals(1, MqttPacketIdGenerator.next(channel));
+ assertEquals(1, MqttPacketIdGenerator.next(otherChannel));
+ assertEquals(2, MqttPacketIdGenerator.next(channel));
+ } finally {
+ MqttPacketIdGenerator.remove(otherChannel);
+ }
+ }
+
+ @Test
+ public void testNextWrapsAroundAfterMaxPacketId() {
+ for (int i = 0; i < MAX_PACKET_ID - 1; i++) {
+ MqttPacketIdGenerator.next(channel);
+ }
+ assertEquals(MAX_PACKET_ID, MqttPacketIdGenerator.next(channel));
+ assertEquals(1, MqttPacketIdGenerator.next(channel));
+ }
+
+ @Test
+ public void testRemoveResetsChannelIdSpace() {
+ assertEquals(1, MqttPacketIdGenerator.next(channel));
+ assertEquals(2, MqttPacketIdGenerator.next(channel));
+ MqttPacketIdGenerator.remove(channel);
+ assertEquals(1, MqttPacketIdGenerator.next(channel));
+ }
+
+}