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-jupiter test + + 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 captor = ArgumentCaptor.forClass(Object.class); + verify(subscriberChannel).writeAndFlush(captor.capture()); + MqttPublishMessage published = (MqttPublishMessage) captor.getValue(); + assertEquals("status/offline", published.variableHeader().topicName()); + assertEquals(1, published.fixedHeader().qosLevel().value()); + assertTrue(published.fixedHeader().isRetain()); + } + + @Test + public void testPublishWillSkipsInactiveChannel() { + when(subscriberChannel.isActive()).thenReturn(false); + when(subscribeRepository.getChannelsByTopic("status/inactive")) + .thenReturn(Collections.singletonMap(subscriberChannel, MqttQoS.AT_LEAST_ONCE)); + + WillRepository.WillEntry will = new WillRepository.WillEntry("status/inactive", "msg".getBytes(), 0, false); + Publish.publishWill(will); + + verify(subscriberChannel, never()).writeAndFlush(any()); + } + + @Test + public void testPublishWillToEmptySubscribers() { + when(subscribeRepository.getChannelsByTopic("topic/none")).thenReturn(Collections.emptyMap()); + + WillRepository.WillEntry will = new WillRepository.WillEntry("topic/none", "msg".getBytes(), 2, false); + Publish.publishWill(will); + } + + @Test + public void testPublishWillQosAndRetain() { + when(subscriberChannel.isActive()).thenReturn(true); + when(subscribeRepository.getChannelsByTopic("qos/retain")) + .thenReturn(Collections.singletonMap(subscriberChannel, MqttQoS.EXACTLY_ONCE)); + + WillRepository.WillEntry will = new WillRepository.WillEntry("qos/retain", "data".getBytes(), 0, false); + Publish.publishWill(will); + + ArgumentCaptor captor = ArgumentCaptor.forClass(Object.class); + verify(subscriberChannel).writeAndFlush(captor.capture()); + MqttPublishMessage published = (MqttPublishMessage) captor.getValue(); + assertEquals(0, published.fixedHeader().qosLevel().value()); + assertFalse(published.fixedHeader().isRetain()); + } + + @Test + public void testPublishWillToWildcardSubscriber() { + when(subscriberChannel.isActive()).thenReturn(true); + when(subscribeRepository.getChannelsByTopic("status/client-001")) + .thenReturn(Collections.singletonMap(subscriberChannel, MqttQoS.AT_MOST_ONCE)); + + WillRepository.WillEntry will = new WillRepository.WillEntry("status/client-001", "gone".getBytes(), 0, false); + Publish.publishWill(will); + + ArgumentCaptor captor = ArgumentCaptor.forClass(Object.class); + verify(subscriberChannel).writeAndFlush(captor.capture()); + MqttPublishMessage published = (MqttPublishMessage) captor.getValue(); + assertEquals("status/client-001", published.variableHeader().topicName()); + } + + @Test + public void testPublishWillCapsQosAtSubscriberGrantedQos() { + when(subscriberChannel.isActive()).thenReturn(true); + when(otherChannel.isActive()).thenReturn(true); + Map subscribers = new HashMap<>(); + subscribers.put(subscriberChannel, MqttQoS.AT_MOST_ONCE); + subscribers.put(otherChannel, MqttQoS.EXACTLY_ONCE); + when(subscribeRepository.getChannelsByTopic("status/qos-cap")).thenReturn(subscribers); + + WillRepository.WillEntry will = new WillRepository.WillEntry("status/qos-cap", "client lost".getBytes(), 2, false); + Publish.publishWill(will); + + ArgumentCaptor captor = ArgumentCaptor.forClass(Object.class); + verify(subscriberChannel).writeAndFlush(captor.capture()); + MqttPublishMessage toQos0Subscriber = (MqttPublishMessage) captor.getValue(); + assertEquals(MqttQoS.AT_MOST_ONCE, toQos0Subscriber.fixedHeader().qosLevel()); + assertEquals(0, toQos0Subscriber.variableHeader().packetId()); + + ArgumentCaptor otherCaptor = ArgumentCaptor.forClass(Object.class); + verify(otherChannel).writeAndFlush(otherCaptor.capture()); + MqttPublishMessage toQos2Subscriber = (MqttPublishMessage) otherCaptor.getValue(); + assertEquals(MqttQoS.EXACTLY_ONCE, toQos2Subscriber.fixedHeader().qosLevel()); + assertTrue(toQos2Subscriber.variableHeader().packetId() > 0); + } +} diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/TopicMatcherTest.java b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/TopicMatcherTest.java new file mode 100644 index 000000000000..561dc65b1f9e --- /dev/null +++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/TopicMatcherTest.java @@ -0,0 +1,100 @@ +/* + * 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 org.junit.jupiter.api.Test; + +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +class TopicMatcherTest { + + @Test + void testExactMatch() { + assertTrue(TopicMatcher.matches("sport/tennis/player1", "sport/tennis/player1")); + assertFalse(TopicMatcher.matches("sport/tennis/player1", "sport/tennis/player2")); + } + + @Test + void testSingleLevelWildcard() { + assertTrue(TopicMatcher.matches("sport/+/player1", "sport/tennis/player1")); + assertTrue(TopicMatcher.matches("sport/+/player1", "sport/football/player1")); + assertTrue(TopicMatcher.matches("+/tennis/+", "sport/tennis/player1")); + assertFalse(TopicMatcher.matches("sport/+/player1", "sport/tennis/stadium/player1")); + assertFalse(TopicMatcher.matches("sport/+", "sport")); + } + + @Test + void testMultiLevelWildcard() { + assertTrue(TopicMatcher.matches("#", "sport/tennis/player1")); + assertTrue(TopicMatcher.matches("#", "sport")); + assertTrue(TopicMatcher.matches("sport/#", "sport")); + assertTrue(TopicMatcher.matches("sport/#", "sport/tennis")); + assertTrue(TopicMatcher.matches("sport/#", "sport/tennis/player1")); + assertTrue(TopicMatcher.matches("sport/#", "sport/tennis/player1/ranking")); + assertFalse(TopicMatcher.matches("sport#", "sport/tennis")); + assertFalse(TopicMatcher.matches("sport/tennis#", "sport/tennis/player1")); + } + + @Test + void testMixedWildcards() { + assertTrue(TopicMatcher.matches("+/tennis/#", "sport/tennis/player1")); + assertTrue(TopicMatcher.matches("+/tennis/#", "sport/tennis/player1/ranking")); + assertTrue(TopicMatcher.matches("sport/+/#", "sport/tennis/player1")); + assertFalse(TopicMatcher.matches("+/tennis/#", "sport/football/player1")); + } + + @Test + void testDollarTopicNotMatchedByLeadingWildcard() { + assertFalse(TopicMatcher.matches("#", "$SYS/broker/version")); + assertFalse(TopicMatcher.matches("+", "$SYS")); + assertFalse(TopicMatcher.matches("+/b", "$SYS/b")); + assertTrue(TopicMatcher.matches("$SYS/#", "$SYS/broker/version")); + assertTrue(TopicMatcher.matches("$SYS/+", "$SYS/broker")); + } + + @Test + void testNullInput() { + assertFalse(TopicMatcher.matches(null, "topic")); + assertFalse(TopicMatcher.matches("topic", null)); + } + + @Test + void testValidFilter() { + assertTrue(TopicMatcher.isValidFilter("sport/tennis/player1")); + assertTrue(TopicMatcher.isValidFilter("sport/+/player1")); + assertTrue(TopicMatcher.isValidFilter("+")); + assertTrue(TopicMatcher.isValidFilter("#")); + assertTrue(TopicMatcher.isValidFilter("sport/#")); + assertTrue(TopicMatcher.isValidFilter("+/tennis/#")); + } + + @Test + void testInvalidFilter() { + assertFalse(TopicMatcher.isValidFilter("sport#")); + assertFalse(TopicMatcher.isValidFilter("sport/tennis#")); + assertFalse(TopicMatcher.isValidFilter("#/sport")); + assertFalse(TopicMatcher.isValidFilter("sport/#/ranking")); + assertFalse(TopicMatcher.isValidFilter("sport+")); + assertFalse(TopicMatcher.isValidFilter("+tennis")); + assertFalse(TopicMatcher.isValidFilter("sport/+tennis")); + assertFalse(TopicMatcher.isValidFilter("")); + assertFalse(TopicMatcher.isValidFilter(null)); + assertFalse(TopicMatcher.isValidFilter("sport/" + (char) 0)); + } +} 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 77e674f3877e..9a3117ad65ea 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 @@ -27,11 +27,13 @@ import org.junit.jupiter.api.Test; import java.time.Duration; +import java.util.ArrayList; 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.CountDownLatch; import java.util.concurrent.ForkJoinPool; import java.util.concurrent.TimeUnit; @@ -53,12 +55,28 @@ public final class SubscribeRepositoryTest { private static final String OTHER_TOPIC = "test/other-topic"; - private static final List ALL_TOPICS = Arrays.asList(ABSENT_TOPIC, TOPIC, OTHER_TOPIC); + private static final String EXACT_TOPIC = "sport/tennis"; + + private static final String CHILD_TOPIC = "sport/tennis/player1"; + + private static final String SINGLE_LEVEL_FILTER = "sport/+/player1"; + + private static final String MULTI_LEVEL_FILTER = "sport/#"; + + private static final String MATCH_ALL_FILTER = "#"; + + private static final String CONCURRENT_TOPIC = "sport/concurrent"; + + private static final List ALL_TOPICS = Arrays.asList( + ABSENT_TOPIC, TOPIC, OTHER_TOPIC, + EXACT_TOPIC, SINGLE_LEVEL_FILTER, MULTI_LEVEL_FILTER, MATCH_ALL_FILTER, CONCURRENT_TOPIC); private static final Duration TIMEOUT = Duration.ofSeconds(5); private static final Duration POLL_INTERVAL = Duration.ofMillis(10); + private static final long JOIN_TIMEOUT_MILLIS = 5000L; + private SubscribeRepository repository; private Channel channel; @@ -197,6 +215,86 @@ public void testRemoveAbsentTopicDoesNotThrow() { assertEquals(MqttQoS.AT_MOST_ONCE, repository.get(TOPIC).get(channel)); } + @Test + public void testGetChannelsByTopicExactMatch() { + repository.add(channel, Collections.singletonList(new MqttTopicSubscription(EXACT_TOPIC, MqttQoS.AT_MOST_ONCE))); + awaitAssert(() -> assertTrue(repository.getChannelsByTopic(EXACT_TOPIC).containsKey(channel))); + assertTrue(repository.getChannelsByTopic(CHILD_TOPIC).isEmpty()); + } + + @Test + public void testGetChannelsByTopicWildcardMatch() { + repository.add(channel, Collections.singletonList(new MqttTopicSubscription(SINGLE_LEVEL_FILTER, MqttQoS.AT_MOST_ONCE))); + awaitAssert(() -> assertTrue(repository.getChannelsByTopic(CHILD_TOPIC).containsKey(channel))); + assertTrue(repository.getChannelsByTopic(EXACT_TOPIC).isEmpty()); + } + + @Test + public void testGetChannelsByTopicDeduplicatesOverlappingSubscriptions() { + repository.add(channel, Arrays.asList( + new MqttTopicSubscription(MATCH_ALL_FILTER, MqttQoS.AT_MOST_ONCE), + new MqttTopicSubscription(MULTI_LEVEL_FILTER, MqttQoS.EXACTLY_ONCE))); + awaitAssert(() -> { + assertTrue(repository.get(MATCH_ALL_FILTER).containsKey(channel)); + assertTrue(repository.get(MULTI_LEVEL_FILTER).containsKey(channel)); + }); + + // MQTT requires at most one delivery per publish per client, at the maximum qos of the matching filters + Map matched = repository.getChannelsByTopic(EXACT_TOPIC); + assertEquals(1, matched.size()); + assertEquals(MqttQoS.EXACTLY_ONCE, matched.get(channel)); + } + + @Test + public void testGetChannelsByTopicMultipleSubscribers() { + repository.add(channel, Collections.singletonList(new MqttTopicSubscription(EXACT_TOPIC, MqttQoS.AT_MOST_ONCE))); + repository.add(otherChannel, Collections.singletonList(new MqttTopicSubscription(EXACT_TOPIC, MqttQoS.AT_MOST_ONCE))); + awaitAssert(() -> assertEquals(2, repository.getChannelsByTopic(EXACT_TOPIC).size())); + assertTrue(repository.getChannelsByTopic(EXACT_TOPIC).keySet().containsAll(Arrays.asList(channel, otherChannel))); + } + + /** + * Regression guard: subscribing to a brand new topic must keep every client + * when several clients send SUBSCRIBE at the same time. + * + * @throws InterruptedException if a subscribing thread is interrupted + */ + @Test + public void concurrentSubscribersOfTheSameNewTopicAreAllRegistered() throws InterruptedException { + int subscriberCount = 8; + CountDownLatch startGate = new CountDownLatch(1); + List subscribers = new ArrayList<>(); + List subscribingThreads = new ArrayList<>(); + for (int i = 0; i < subscriberCount; i++) { + Channel subscriber = mock(Channel.class); + subscribers.add(subscriber); + Thread thread = new Thread(() -> subscribeAfter(startGate, subscriber, CONCURRENT_TOPIC)); + thread.start(); + subscribingThreads.add(thread); + } + + startGate.countDown(); + for (Thread thread : subscribingThreads) { + thread.join(JOIN_TIMEOUT_MILLIS); + assertFalse(thread.isAlive(), "subscribing thread did not finish in time"); + } + + awaitAssert(() -> assertEquals(subscriberCount, repository.get(CONCURRENT_TOPIC).size())); + Map matched = repository.getChannelsByTopic(CONCURRENT_TOPIC); + assertEquals(subscriberCount, matched.size()); + assertTrue(matched.keySet().containsAll(subscribers)); + } + + private void subscribeAfter(final CountDownLatch startGate, final Channel subscriber, final String topic) { + try { + startGate.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + return; + } + repository.add(subscriber, Collections.singletonList(new MqttTopicSubscription(topic, MqttQoS.AT_MOST_ONCE))); + } + /** * The repository mutates its state asynchronously on the common pool, * so assertions have to be retried until the mutation becomes visible. diff --git a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/WillRepositoryTest.java b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/WillRepositoryTest.java new file mode 100644 index 000000000000..5692311a3c9d --- /dev/null +++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/WillRepositoryTest.java @@ -0,0 +1,112 @@ +/* + * 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.embedded.EmbeddedChannel; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +import static org.hamcrest.MatcherAssert.assertThat; +import static org.hamcrest.Matchers.notNullValue; +import static org.hamcrest.Matchers.nullValue; +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.assertTrue; + +public class WillRepositoryTest { + + private WillRepository willRepository; + + private EmbeddedChannel channel; + + @BeforeEach + public void setUp() { + willRepository = new WillRepository(); + channel = new EmbeddedChannel(); + } + + @AfterEach + public void tearDown() { + channel.close(); + } + + @Test + public void testAddAndGetWillEntry() { + byte[] message = "offline".getBytes(); + WillRepository.WillEntry entry = new WillRepository.WillEntry("topic/test", message, 1, true); + willRepository.add(channel, entry); + + WillRepository.WillEntry result = willRepository.get(channel); + assertThat(result, notNullValue()); + assertEquals("topic/test", result.getTopic()); + assertArrayEquals(message, result.getMessage()); + assertEquals(1, result.getQos()); + assertTrue(result.isRetain()); + } + + @Test + public void testRemoveWillEntry() { + WillRepository.WillEntry entry = new WillRepository.WillEntry("topic/test", "bye".getBytes(), 0, false); + willRepository.add(channel, entry); + assertThat(willRepository.get(channel), notNullValue()); + + willRepository.remove(channel); + assertThat(willRepository.get(channel), nullValue()); + } + + @Test + public void testGetReturnsNullForUnknownChannel() { + assertThat(willRepository.get(channel), nullValue()); + } + + @Test + public void testWillEntryFields() { + byte[] message = "disconnected".getBytes(); + WillRepository.WillEntry entry = new WillRepository.WillEntry("alerts/disconnect", message, 2, false); + + assertEquals("alerts/disconnect", entry.getTopic()); + assertArrayEquals("disconnected".getBytes(), entry.getMessage()); + assertEquals(2, entry.getQos()); + assertFalse(entry.isRetain()); + } + + @Test + public void testReplaceWillEntryOnReconnect() { + WillRepository.WillEntry first = new WillRepository.WillEntry("topic/a", "msg1".getBytes(), 0, false); + WillRepository.WillEntry second = new WillRepository.WillEntry("topic/b", "msg2".getBytes(), 1, true); + willRepository.add(channel, first); + willRepository.add(channel, second); + + WillRepository.WillEntry result = willRepository.get(channel); + assertEquals("topic/b", result.getTopic()); + assertArrayEquals("msg2".getBytes(), result.getMessage()); + } + + @Test + public void testRemoveNonexistentChannel() { + EmbeddedChannel otherChannel = new EmbeddedChannel(); + try { + willRepository.remove(otherChannel); + assertThat(willRepository.get(channel), nullValue()); + } finally { + otherChannel.close(); + } + } +}