From 942c54704db575d5a96b339934bde26ce869051b Mon Sep 17 00:00:00 2001 From: wy471x Date: Mon, 10 Aug 2026 00:08:59 +0800 Subject: [PATCH 1/6] feat: implement MQTT Last Will and Testament (LWT) Parse and store will fields (topic, message, QoS, retain) in WillRepository on CONNECT. Publish will message to subscribers on ungraceful disconnect via channelInactive hook. Clear will on graceful DISCONNECT. Fix DISCONNECT message dispatch in MqttFactory that was previously dropped. Co-Authored-By: Claude Opus 4.7 --- shenyu-protocol/shenyu-protocol-mqtt/pom.xml | 33 ++++ .../apache/shenyu/protocol/mqtt/Connect.java | 12 ++ .../shenyu/protocol/mqtt/Disconnect.java | 4 +- .../shenyu/protocol/mqtt/MqttFactory.java | 4 +- .../protocol/mqtt/MqttTransportHandler.java | 14 ++ .../apache/shenyu/protocol/mqtt/Publish.java | 21 +++ .../mqtt/repositories/WillRepository.java | 85 +++++++++ .../shenyu/protocol/mqtt/ConnectTest.java | 162 ++++++++++++++++++ .../shenyu/protocol/mqtt/DisconnectTest.java | 89 ++++++++++ .../mqtt/MqttTransportHandlerTest.java | 101 +++++++++++ .../shenyu/protocol/mqtt/PublishWillTest.java | 115 +++++++++++++ .../mqtt/repositories/WillRepositoryTest.java | 112 ++++++++++++ 12 files changed, 749 insertions(+), 3 deletions(-) create mode 100644 shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/WillRepository.java create mode 100644 shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java create mode 100644 shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/DisconnectTest.java create mode 100644 shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java create mode 100644 shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishWillTest.java create mode 100644 shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/WillRepositoryTest.java diff --git a/shenyu-protocol/shenyu-protocol-mqtt/pom.xml b/shenyu-protocol/shenyu-protocol-mqtt/pom.xml index 2ddc80e0afdd..f0cb4fe4589c 100644 --- a/shenyu-protocol/shenyu-protocol-mqtt/pom.xml +++ b/shenyu-protocol/shenyu-protocol-mqtt/pom.xml @@ -46,6 +46,39 @@ reflections 0.9.11 + + + org.junit.jupiter + junit-jupiter + test + + + org.mockito + mockito-junit-jupiter + ${mockito.version} + test + + + org.mockito + mockito-core + ${mockito.version} + test + + + + + org.apache.maven.plugins + maven-surefire-plugin + + --add-opens java.base/java.lang=ALL-UNNAMED + --add-opens java.base/java.util=ALL-UNNAMED + --add-opens java.base/java.net=ALL-UNNAMED + -Dnet.bytebuddy.experimental=true + + + + + 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 8a9fa005a41e..b59e8d6fb995 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; @@ -67,6 +68,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 87a06e39d2b3..a6a186c3bdae 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.ChannelHandlerContext; import org.apache.shenyu.common.utils.Singleton; import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository; +import org.apache.shenyu.protocol.mqtt.repositories.WillRepository; /** * The DISCONNECT message is sent from the client to the server to indicate @@ -36,8 +37,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 9a59f1407ee5..e87b2d6cbccb 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,8 +65,10 @@ public void connect() { case PINGREQ: messageType.pingReq(ctx); 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 49ddc8a058a3..ebca15a90a71 100644 --- a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java +++ b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java @@ -22,6 +22,10 @@ import io.netty.handler.codec.mqtt.MqttMessage; import io.netty.util.concurrent.Future; import io.netty.util.concurrent.GenericFutureListener; +import org.apache.shenyu.common.utils.Singleton; +import org.apache.shenyu.protocol.mqtt.repositories.WillRepository; + +import java.util.Objects; /** * mqtt transport handler. @@ -38,6 +42,16 @@ public void channelRead(final ChannelHandlerContext ctx, final Object msg) throw } } + @Override + public void channelInactive(final ChannelHandlerContext ctx) throws Exception { + WillRepository.WillEntry will = Singleton.INST.get(WillRepository.class).get(ctx.channel()); + if (Objects.nonNull(will)) { + Publish.publishWill(will); + Singleton.INST.get(WillRepository.class).remove(ctx.channel()); + } + super.channelInactive(ctx); + } + @Override public void operationComplete(final Future future) throws Exception { 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 b8804cf6fa4e..aaad7f6cce66 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 @@ -32,6 +32,8 @@ import org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository; import org.apache.shenyu.protocol.mqtt.repositories.TopicRepository; +import org.apache.shenyu.protocol.mqtt.repositories.WillRepository; + import java.util.List; import java.util.concurrent.CompletableFuture; @@ -124,4 +126,23 @@ private void send(final String topic, final ByteBuf payload, final int packetId) } }); } + + /** + * 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) { + List channels = Singleton.INST.get(SubscribeRepository.class).get(will.getTopic()); + MqttQoS willQos = MqttQoS.valueOf(will.getQos()); + channels.parallelStream().forEach(channel -> { + if (channel.isActive()) { + MqttFixedHeader mqttFixedHeader = new MqttFixedHeader(MqttMessageType.PUBLISH, false, willQos, will.isRetain(), 0); + MqttPublishVariableHeader mqttPublishVariableHeader = new MqttPublishVariableHeader(will.getTopic(), 0); + 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/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 new file mode 100644 index 000000000000..8f561e1f64b9 --- /dev/null +++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java @@ -0,0 +1,162 @@ +/* + * 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 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.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.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.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; +import static org.mockito.Mockito.when; + +@ExtendWith(MockitoExtension.class) +public class ConnectTest { + + private static final String VALID_USER = "admin"; + + private static final String VALID_PASS = "pass123"; + + @Mock + private ChannelHandlerContext ctx; + + @Mock + private Channel channel; + + @Mock + private MqttConnectMessage msg; + + @Mock + private MqttConnectVariableHeader variableHeader; + + @Mock + private MqttConnectPayload payload; + + private Connect connect; + + private WillRepository willRepository; + + private ChannelRepository channelRepository; + + @BeforeEach + public void setUp() { + connect = new Connect(); + willRepository = new WillRepository(); + channelRepository = new ChannelRepository(); + Singleton.INST.single(WillRepository.class, willRepository); + Singleton.INST.single(ChannelRepository.class, channelRepository); + + when(ctx.channel()).thenReturn(channel); + when(msg.variableHeader()).thenReturn(variableHeader); + when(msg.payload()).thenReturn(payload); + when(variableHeader.version()).thenReturn((int) MqttVersion.MQTT_3_1.protocolLevel()); + when(payload.clientIdentifier()).thenReturn("test-client-001"); + when(payload.userName()).thenReturn(VALID_USER); + when(payload.passwordInBytes()).thenReturn(VALID_PASS.getBytes()); + + MqttContext mqttContext = new MqttContext(); + mqttContext.setUserName(VALID_USER); + mqttContext.setPassword(VALID_PASS); + } + + @AfterEach + public void tearDown() { + Singleton.INST.single(WillRepository.class, new WillRepository()); + Singleton.INST.single(ChannelRepository.class, new ChannelRepository()); + } + + @Test + public void testStoresWillOnConnect() { + byte[] willMessage = "client disconnected unexpectedly".getBytes(); + when(variableHeader.isWillFlag()).thenReturn(true); + when(variableHeader.willQos()).thenReturn(1); + when(variableHeader.isWillRetain()).thenReturn(true); + when(payload.willTopic()).thenReturn("status/client-001"); + when(payload.willMessageInBytes()).thenReturn(willMessage); + + connect.connect(ctx, msg); + + WillRepository.WillEntry will = willRepository.get(channel); + assertThat(will, notNullValue()); + assertEquals("status/client-001", will.getTopic()); + assertArrayEquals(willMessage, will.getMessage()); + assertEquals(1, will.getQos()); + assertTrue(will.isRetain()); + } + + @Test + public void testDoesNotStoreWillWhenWillFlagIsFalse() { + when(variableHeader.isWillFlag()).thenReturn(false); + + connect.connect(ctx, msg); + + WillRepository.WillEntry will = willRepository.get(channel); + assertThat(will, nullValue()); + } + + @Test + public void testWillQosZero() { + byte[] willMessage = "qos0 will".getBytes(); + when(variableHeader.isWillFlag()).thenReturn(true); + when(variableHeader.willQos()).thenReturn(0); + when(variableHeader.isWillRetain()).thenReturn(false); + when(payload.willTopic()).thenReturn("topic/qos0"); + when(payload.willMessageInBytes()).thenReturn(willMessage); + + connect.connect(ctx, msg); + + WillRepository.WillEntry will = willRepository.get(channel); + assertThat(will, notNullValue()); + assertEquals(0, will.getQos()); + assertFalse(will.isRetain()); + } + + @Test + public void testWillRetainTrue() { + byte[] willMessage = "retained will".getBytes(); + when(variableHeader.isWillFlag()).thenReturn(true); + when(variableHeader.willQos()).thenReturn(2); + when(variableHeader.isWillRetain()).thenReturn(true); + when(payload.willTopic()).thenReturn("topic/retained"); + when(payload.willMessageInBytes()).thenReturn(willMessage); + + connect.connect(ctx, msg); + + WillRepository.WillEntry will = willRepository.get(channel); + assertThat(will, notNullValue()); + assertTrue(will.isRetain()); + assertEquals(2, will.getQos()); + } +} 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/MqttTransportHandlerTest.java b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java new file mode 100644 index 000000000000..5bc26bb3b4d0 --- /dev/null +++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java @@ -0,0 +1,101 @@ +/* + * 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.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.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 MqttTransportHandlerTest { + + @Mock + private ChannelHandlerContext ctx; + + @Mock + private Channel channel; + + @Mock + private SubscribeRepository subscribeRepository; + + private MqttTransportHandler handler; + + private WillRepository willRepository; + + @BeforeEach + public void setUp() { + handler = new MqttTransportHandler(); + willRepository = new WillRepository(); + Singleton.INST.single(WillRepository.class, willRepository); + Singleton.INST.single(SubscribeRepository.class, subscribeRepository); + when(ctx.channel()).thenReturn(channel); + } + + @AfterEach + public void tearDown() { + Singleton.INST.single(WillRepository.class, new WillRepository()); + Singleton.INST.single(SubscribeRepository.class, new SubscribeRepository()); + } + + @Test + public void testChannelInactiveFiresWillAndRemovesIt() throws Exception { + byte[] willMessage = "sudden disconnect".getBytes(); + WillRepository.WillEntry will = new WillRepository.WillEntry("status/offline", willMessage, 1, true); + willRepository.add(channel, will); + + // publishWill uses subscribeRepository to get target channels + when(subscribeRepository.get("status/offline")).thenReturn(java.util.Collections.emptyList()); + + handler.channelInactive(ctx); + + // will should be removed after firing + assertThat(willRepository.get(channel), nullValue()); + } + + @Test + public void testChannelInactiveDoesNothingWhenNoWill() throws Exception { + handler.channelInactive(ctx); + + assertThat(willRepository.get(channel), nullValue()); + } + + @Test + public void testChannelInactiveAfterDisconnectClearsWill() throws Exception { + byte[] willMessage = "graceful close".getBytes(); + WillRepository.WillEntry will = new WillRepository.WillEntry("status/clean", willMessage, 0, false); + willRepository.add(channel, will); + + // simulate graceful disconnect: remove will first, then channelInactive + willRepository.remove(channel); + handler.channelInactive(ctx); + + assertThat(willRepository.get(channel), nullValue()); + } +} 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..e5be0f2c9fd9 --- /dev/null +++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishWillTest.java @@ -0,0 +1,115 @@ +/* + * 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 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 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; + + @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.get("status/offline")) + .thenReturn(Collections.singletonList(subscriberChannel)); + + 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.get("status/inactive")) + .thenReturn(Collections.singletonList(subscriberChannel)); + + 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.get("topic/none")).thenReturn(Collections.emptyList()); + + 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.get("qos/retain")) + .thenReturn(Collections.singletonList(subscriberChannel)); + + 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()); + } +} 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(); + } + } +} From ab9592154d07d428debe1fee7d27af6b88550c6e Mon Sep 17 00:00:00 2001 From: xiaoyu <549477611@qq.com> Date: Tue, 11 Aug 2026 18:22:10 +0800 Subject: [PATCH 2/6] Potential fix for pull request finding Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- .../org/apache/shenyu/protocol/mqtt/Publish.java | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) 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 aaad7f6cce66..154f6580c066 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 @@ -133,12 +133,18 @@ private void send(final String topic, final ByteBuf payload, final int packetId) * @param will the will entry containing topic, message, qos, and retain flag */ static void publishWill(final WillRepository.WillEntry will) { - List channels = Singleton.INST.get(SubscribeRepository.class).get(will.getTopic()); - MqttQoS willQos = MqttQoS.valueOf(will.getQos()); + if (will == null || will.getTopic() == null || will.getMessage() == null) { + return; + } + final List channels = Singleton.INST.get(SubscribeRepository.class).get(will.getTopic()); + final MqttQoS willQos = MqttQoS.valueOf(will.getQos()); + final int packetId = willQos == MqttQoS.AT_MOST_ONCE + ? 0 + : java.util.concurrent.ThreadLocalRandom.current().nextInt(1, 65536); channels.parallelStream().forEach(channel -> { if (channel.isActive()) { MqttFixedHeader mqttFixedHeader = new MqttFixedHeader(MqttMessageType.PUBLISH, false, willQos, will.isRetain(), 0); - MqttPublishVariableHeader mqttPublishVariableHeader = new MqttPublishVariableHeader(will.getTopic(), 0); + MqttPublishVariableHeader mqttPublishVariableHeader = new MqttPublishVariableHeader(will.getTopic(), packetId); MqttPublishMessage mqttPublishMessage = new MqttPublishMessage(mqttFixedHeader, mqttPublishVariableHeader, Unpooled.wrappedBuffer(will.getMessage())); channel.writeAndFlush(mqttPublishMessage); From 61384ab108e0ff55a80a21af54e633705f8ca273 Mon Sep 17 00:00:00 2001 From: wy471x Date: Sat, 15 Aug 2026 23:58:19 +0800 Subject: [PATCH 3/6] feat: route will delivery through wildcard topic matching Publish.publishWill used an exact SubscribeRepository.get() lookup, so clients subscribed to wildcard filters (e.g. status/#) never received wills published to concrete topics like status/client-001. Port the TopicMatcher and SubscribeRepository.getChannelsByTopic from #6906 and route will delivery through it, consistent with normal publish routing. Co-Authored-By: Claude Opus 4.7 --- .../apache/shenyu/protocol/mqtt/Publish.java | 5 +- .../shenyu/protocol/mqtt/TopicMatcher.java | 103 +++++++++++++++ .../repositories/SubscribeRepository.java | 34 +++++ .../mqtt/MqttTransportHandlerTest.java | 2 +- .../shenyu/protocol/mqtt/PublishWillTest.java | 23 +++- .../protocol/mqtt/TopicMatcherTest.java | 100 +++++++++++++++ .../repositories/SubscribeRepositoryTest.java | 117 ++++++++++++++++++ 7 files changed, 377 insertions(+), 7 deletions(-) create mode 100644 shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/TopicMatcher.java create mode 100644 shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/TopicMatcherTest.java create mode 100644 shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepositoryTest.java 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 154f6580c066..2867322b4555 100644 --- a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java +++ b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java @@ -35,6 +35,7 @@ import org.apache.shenyu.protocol.mqtt.repositories.WillRepository; import java.util.List; +import java.util.Objects; import java.util.concurrent.CompletableFuture; import static io.netty.handler.codec.mqtt.MqttMessageType.PUBACK; @@ -133,10 +134,10 @@ private void send(final String topic, final ByteBuf payload, final int packetId) * @param will the will entry containing topic, message, qos, and retain flag */ static void publishWill(final WillRepository.WillEntry will) { - if (will == null || will.getTopic() == null || will.getMessage() == null) { + if (Objects.isNull(will) || Objects.isNull(will.getTopic()) || Objects.isNull(will.getMessage())) { return; } - final List channels = Singleton.INST.get(SubscribeRepository.class).get(will.getTopic()); + final List channels = Singleton.INST.get(SubscribeRepository.class).getChannelsByTopic(will.getTopic()); final MqttQoS willQos = MqttQoS.valueOf(will.getQos()); final int packetId = willQos == MqttQoS.AT_MOST_ONCE ? 0 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 39e82c44009e..ad3d347193f6 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 @@ -19,11 +19,15 @@ import io.netty.channel.Channel; import io.netty.handler.codec.mqtt.MqttTopicSubscription; +import org.apache.shenyu.protocol.mqtt.TopicMatcher; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import java.util.ArrayList; +import java.util.LinkedHashSet; import java.util.List; import java.util.Map; +import java.util.Objects; import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; @@ -91,4 +95,34 @@ public List get(final String topic) { return TOPIC_CHANNEL_FACTORY.getOrDefault(topic, new CopyOnWriteArrayList<>()); } + /** + * Get channels whose subscription filter matches the published topic. + * Supports MQTT wildcards: + (single-level) and # (multi-level). + * + * @param topic the published topic name + * @return channels subscribed to matching topic filters + */ + public List getChannelsByTopic(final String topic) { + // MQTT requires at most one delivery per publish per client, so dedupe + // channels when overlapping filters (e.g. sport/# and #) both match. + Set result = new LinkedHashSet<>(); + + // fast path: exact subscription, no wildcard scan needed + List exactMatch = TOPIC_CHANNEL_FACTORY.get(topic); + if (Objects.nonNull(exactMatch)) { + result.addAll(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)) { + result.addAll(entry.getValue()); + } + } + return new ArrayList<>(result); + } + } 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 5bc26bb3b4d0..7ca65e51e366 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 @@ -71,7 +71,7 @@ public void testChannelInactiveFiresWillAndRemovesIt() throws Exception { willRepository.add(channel, will); // publishWill uses subscribeRepository to get target channels - when(subscribeRepository.get("status/offline")).thenReturn(java.util.Collections.emptyList()); + when(subscribeRepository.getChannelsByTopic("status/offline")).thenReturn(java.util.Collections.emptyList()); handler.channelInactive(ctx); 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 index e5be0f2c9fd9..7ba6d1590c7d 100644 --- 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 @@ -62,7 +62,7 @@ public void tearDown() { @Test public void testPublishWillToActiveSubscriber() { when(subscriberChannel.isActive()).thenReturn(true); - when(subscribeRepository.get("status/offline")) + when(subscribeRepository.getChannelsByTopic("status/offline")) .thenReturn(Collections.singletonList(subscriberChannel)); byte[] message = "client lost".getBytes(); @@ -80,7 +80,7 @@ public void testPublishWillToActiveSubscriber() { @Test public void testPublishWillSkipsInactiveChannel() { when(subscriberChannel.isActive()).thenReturn(false); - when(subscribeRepository.get("status/inactive")) + when(subscribeRepository.getChannelsByTopic("status/inactive")) .thenReturn(Collections.singletonList(subscriberChannel)); WillRepository.WillEntry will = new WillRepository.WillEntry("status/inactive", "msg".getBytes(), 0, false); @@ -91,7 +91,7 @@ public void testPublishWillSkipsInactiveChannel() { @Test public void testPublishWillToEmptySubscribers() { - when(subscribeRepository.get("topic/none")).thenReturn(Collections.emptyList()); + when(subscribeRepository.getChannelsByTopic("topic/none")).thenReturn(Collections.emptyList()); WillRepository.WillEntry will = new WillRepository.WillEntry("topic/none", "msg".getBytes(), 2, false); Publish.publishWill(will); @@ -100,7 +100,7 @@ public void testPublishWillToEmptySubscribers() { @Test public void testPublishWillQosAndRetain() { when(subscriberChannel.isActive()).thenReturn(true); - when(subscribeRepository.get("qos/retain")) + when(subscribeRepository.getChannelsByTopic("qos/retain")) .thenReturn(Collections.singletonList(subscriberChannel)); WillRepository.WillEntry will = new WillRepository.WillEntry("qos/retain", "data".getBytes(), 0, false); @@ -112,4 +112,19 @@ public void testPublishWillQosAndRetain() { 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.singletonList(subscriberChannel)); + + 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()); + } } 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 new file mode 100644 index 000000000000..824f7702ab59 --- /dev/null +++ b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepositoryTest.java @@ -0,0 +1,117 @@ +/* + * 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 io.netty.channel.embedded.EmbeddedChannel; +import io.netty.handler.codec.mqtt.MqttQoS; +import io.netty.handler.codec.mqtt.MqttTopicSubscription; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.function.BooleanSupplier; +import java.util.stream.Collectors; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.fail; + +class SubscribeRepositoryTest { + + private final SubscribeRepository repository = new SubscribeRepository(); + + private final List channels = new ArrayList<>(); + + private final List topicFilters = new ArrayList<>(); + + @AfterEach + void cleanup() throws InterruptedException { + if (!topicFilters.isEmpty()) { + for (Channel subscribed : channels) { + repository.remove(topicFilters, subscribed); + } + awaitUntil(() -> topicFilters.stream().allMatch(topic -> repository.get(topic).isEmpty())); + } + channels.forEach(channel -> ((EmbeddedChannel) channel).finishAndReleaseAll()); + } + + @Test + void testGetChannelsByTopicExactMatch() throws InterruptedException { + Channel subscriber = newSubscriber("sport/tennis"); + + assertTrue(repository.getChannelsByTopic("sport/tennis").contains(subscriber)); + assertTrue(repository.getChannelsByTopic("sport/tennis/player1").isEmpty()); + } + + @Test + void testGetChannelsByTopicWildcardMatch() throws InterruptedException { + Channel subscriber = newSubscriber("sport/+/player1"); + + assertTrue(repository.getChannelsByTopic("sport/tennis/player1").contains(subscriber)); + assertTrue(repository.getChannelsByTopic("sport/tennis").isEmpty()); + } + + @Test + void testGetChannelsByTopicDeduplicatesOverlappingSubscriptions() throws InterruptedException { + Channel subscriber = newSubscriber("#", "sport/#"); + + // MQTT requires at most one delivery per publish per client + assertEquals(1, repository.getChannelsByTopic("sport/tennis").size()); + assertTrue(repository.getChannelsByTopic("sport/tennis").contains(subscriber)); + } + + @Test + void testGetChannelsByTopicMultipleSubscribers() throws InterruptedException { + final Channel first = newSubscriber("sport/tennis"); + final Channel second = new EmbeddedChannel(); + channels.add(second); + topicFilters.add("sport/tennis"); + repository.add(second, subscriptions("sport/tennis")); + awaitUntil(() -> repository.get("sport/tennis").containsAll(Arrays.asList(first, second))); + + assertEquals(2, repository.getChannelsByTopic("sport/tennis").size()); + } + + private Channel newSubscriber(final String... topics) throws InterruptedException { + Channel subscriber = new EmbeddedChannel(); + channels.add(subscriber); + topicFilters.addAll(Arrays.asList(topics)); + repository.add(subscriber, subscriptions(topics)); + awaitUntil(() -> Arrays.stream(topics).allMatch(topic -> repository.get(topic).contains(subscriber))); + return subscriber; + } + + private List subscriptions(final String... topics) { + return Arrays.stream(topics) + .map(topic -> new MqttTopicSubscription(topic, MqttQoS.AT_MOST_ONCE)) + .collect(Collectors.toList()); + } + + private void awaitUntil(final BooleanSupplier condition) throws InterruptedException { + long deadline = System.currentTimeMillis() + 5000; + while (!condition.getAsBoolean()) { + if (System.currentTimeMillis() >= deadline) { + fail("condition not met within timeout"); + } + Thread.sleep(10); + } + } +} From 2a8a5d9c8589b63e535f8b2970a83d5dafc102ac Mon Sep 17 00:00:00 2001 From: wy471x Date: Sat, 19 Sep 2026 20:24:02 +0800 Subject: [PATCH 4/6] fix(mqtt): make SubscribeRepository updates synchronous and race-free - allocate the per-topic channel list with computeIfAbsent and apply add/remove on the calling thread, so concurrent subscribers of the same new topic no longer overwrite each other's list and a subscription is visible as soon as add returns; get(List) no longer throws when a topic has no subscribers; drop the unused logger - SubscribeRepositoryTest: replace the Mockito mock and the awaitility/common-pool polling with deterministic assertions, and add a concurrency regression test for eight subscribers of the same new topic - ConnectTest: drive Connect with real MQTT messages instead of mocks, keep the credentials fixed for the whole class, cover the identifier-rejected and bad-credentials branches and assert a rejected CONNECT stores no will - MqttFactoryTest: DISCONNECT is dispatched to Disconnect, so assert the channel is closed and the channel/will repositories are cleared, and register the WillRepository that Disconnect looks up via Singleton - MqttTransportHandler: drop the will before publishing it and notify the pipeline exactly once on channelInactive --- .../protocol/mqtt/MqttTransportHandler.java | 13 +- .../repositories/SubscribeRepository.java | 39 ++- .../shenyu/protocol/mqtt/ConnectTest.java | 286 ++++++++++-------- .../shenyu/protocol/mqtt/MqttFactoryTest.java | 23 +- .../mqtt/MqttTransportHandlerTest.java | 32 +- .../repositories/SubscribeRepositoryTest.java | 117 +++---- 6 files changed, 294 insertions(+), 216 deletions(-) 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 61f979cc9ccb..1a849456aba0 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,7 @@ package org.apache.shenyu.protocol.mqtt; +import io.netty.channel.Channel; import io.netty.channel.ChannelHandlerContext; import io.netty.channel.ChannelInboundHandlerAdapter; import io.netty.handler.codec.mqtt.MqttMessage; @@ -45,14 +46,18 @@ 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); - WillRepository.WillEntry will = Singleton.INST.get(WillRepository.class).get(ctx.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); - Singleton.INST.get(WillRepository.class).remove(ctx.channel()); } + // local state is consistent now, notify the rest of the pipeline exactly once. super.channelInactive(ctx); } 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 6846316033a1..05b1f05fefe3 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 @@ -21,36 +21,33 @@ import io.netty.handler.codec.mqtt.MqttTopicSubscription; import org.apache.commons.collections4.CollectionUtils; import org.apache.shenyu.protocol.mqtt.TopicMatcher; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; import java.util.ArrayList; +import java.util.Collections; import java.util.LinkedHashSet; import java.util.List; import java.util.Map; import java.util.Objects; import java.util.Set; -import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.CopyOnWriteArrayList; -import java.util.concurrent.CopyOnWriteArraySet; /** * Topic and channel association. + * + *

Subscription updates are applied synchronously on the calling (event loop) thread and every + * topic holds a copy-on-write list of channels, so a subscription is visible to publish as soon as + * {@code add} returns and concurrent subscribers of the same topic never overwrite each other. */ public class SubscribeRepository implements BaseRepository, List> { - private static final Logger LOG = LoggerFactory.getLogger(SubscribeRepository.class); - private static final Map> TOPIC_CHANNEL_FACTORY = new ConcurrentHashMap<>(); @Override public void add(final List topics, final List channels) { - CompletableFuture.runAsync(() -> topics.parallelStream().forEach(s -> { - List list = get(s); - list.addAll(channels); - TOPIC_CHANNEL_FACTORY.put(s, list); - })); + topics.forEach(topic -> TOPIC_CHANNEL_FACTORY + .computeIfAbsent(topic, key -> new CopyOnWriteArrayList<>()) + .addAll(channels)); } /** @@ -59,16 +56,14 @@ public void add(final List topics, final List channels) { * @param mqttTopicSubscription mqtt subscription info */ public void add(final Channel channel, final List mqttTopicSubscription) { - CompletableFuture.runAsync(() -> mqttTopicSubscription.parallelStream().forEach(s -> { - List channels = get(s.topicName()); - channels.add(channel); - TOPIC_CHANNEL_FACTORY.put(s.topicName(), channels); - })); + mqttTopicSubscription.forEach(subscription -> TOPIC_CHANNEL_FACTORY + .computeIfAbsent(subscription.topicName(), key -> new CopyOnWriteArrayList<>()) + .add(channel)); } @Override public void remove(final List topics) { - CompletableFuture.runAsync(() -> topics.parallelStream().forEach(TOPIC_CHANNEL_FACTORY::remove)); + topics.forEach(TOPIC_CHANNEL_FACTORY::remove); } /** @@ -77,19 +72,19 @@ public void remove(final List topics) { * @param channel channel */ public void remove(final List topics, final Channel channel) { - CompletableFuture.runAsync(() -> topics.parallelStream().forEach(topic -> { + topics.forEach(topic -> { List channels = TOPIC_CHANNEL_FACTORY.get(topic); if (CollectionUtils.isNotEmpty(channels)) { channels.remove(channel); } - })); + }); } @Override public List get(final List topics) { - Set channels = new CopyOnWriteArraySet<>(); - topics.parallelStream().forEach(s -> channels.addAll(TOPIC_CHANNEL_FACTORY.get(s))); - return new CopyOnWriteArrayList<>(channels); + Set channels = new LinkedHashSet<>(); + topics.forEach(topic -> channels.addAll(TOPIC_CHANNEL_FACTORY.getOrDefault(topic, Collections.emptyList()))); + return new ArrayList<>(channels); } /** 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 7d57dc5fe763..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 @@ -17,7 +17,6 @@ package org.apache.shenyu.protocol.mqtt; -import io.netty.channel.Channel; import io.netty.channel.ChannelHandlerContext; import io.netty.channel.ChannelInboundHandlerAdapter; import io.netty.channel.embedded.EmbeddedChannel; @@ -31,82 +30,83 @@ 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.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 org.junit.jupiter.api.extension.ExtendWith; -import org.mockito.Mock; -import org.mockito.junit.jupiter.MockitoExtension; 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.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.assertNotNull; import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertTrue; -import static org.mockito.Mockito.when; /** * Test cases for {@link Connect}. */ -@ExtendWith(MockitoExtension.class) public final class ConnectTest { - private static final String VALID_USER = "admin"; - - private static final String VALID_PASS = "pass123"; - - 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; - - @Mock - private ChannelHandlerContext ctx; + private static final String CLIENT_ID = "test-client"; - @Mock - private Channel channel; + private static final String WILL_TOPIC = "status/client-001"; - @Mock - private MqttConnectMessage msg; + private final List channels = new ArrayList<>(); - @Mock - private MqttConnectVariableHeader variableHeader; + private ChannelRepository channelRepository; - @Mock - private MqttConnectPayload payload; + private WillRepository willRepository; private Connect connect; - private WillRepository willRepository; - @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 @@ -126,99 +126,70 @@ 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 duplicateConnectIsRejected() { - EmbeddedChannel channel = new EmbeddedChannel(new ChannelInboundHandlerAdapter()); - ChannelHandlerContext ctx = channel.pipeline().lastContext(); + public void emptyClientIdIsRejected() { + EmbeddedChannel channel = newChannel(); - new Connect().connect(ctx, connectMessage(MqttVersion.MQTT_3_1_1.protocolName(), MqttVersion.MQTT_3_1_1.protocolLevel())); - assertNotNull(channel.readOutbound()); - - new Connect().connect(ctx, connectMessage(MqttVersion.MQTT_3_1_1.protocolName(), MqttVersion.MQTT_3_1_1.protocolLevel())); + connect.connect(context(channel), buildConnectMessage(MqttVersion.MQTT_3_1_1.protocolName(), + MqttVersion.MQTT_3_1_1.protocolLevel(), "", PASSWORD, false, 0, false, null, null)); - channel.runPendingTasks(); - assertFalse(channel.isActive()); - assertNull(channel.readOutbound()); + MqttConnAckMessage ackMessage = channel.readOutbound(); + assertNotNull(ackMessage); + assertEquals(CONNECTION_REFUSED_IDENTIFIER_REJECTED, ackMessage.variableHeader().connectReturnCode()); + assertChannelClosed(channel); + assertNull(channelRepository.get(channel)); } - private void connectIsAccepted(final MqttVersion version) { - EmbeddedChannel channel = new EmbeddedChannel(new ChannelInboundHandlerAdapter()); - ChannelHandlerContext ctx = channel.pipeline().lastContext(); + @Test + public void invalidCredentialsAreRejected() { + EmbeddedChannel channel = newChannel(); - new Connect().connect(ctx, connectMessage(version.protocolName(), version.protocolLevel())); + 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_ACCEPTED, ackMessage.variableHeader().connectReturnCode()); - assertTrue(ackMessage.variableHeader().isSessionPresent()); - await().atMost(Duration.ofSeconds(5)) - .until(() -> CLIENT_ID.equals(channelRepository.get(channel))); - } - - private MqttConnectMessage connectMessage(final String protocolName, final int protocolLevel) { - 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)); - return new MqttConnectMessage(fixedHeader, variableHeader, payload); + assertEquals(CONNECTION_REFUSED_BAD_USER_NAME_OR_PASSWORD, ackMessage.variableHeader().connectReturnCode()); + assertChannelClosed(channel); + assertNull(channelRepository.get(channel)); } - @BeforeEach - public void setUp() { - connect = new Connect(); - willRepository = new WillRepository(); - channelRepository = new ChannelRepository(); - Singleton.INST.single(WillRepository.class, willRepository); - Singleton.INST.single(ChannelRepository.class, channelRepository); + @Test + public void duplicateConnectIsRejected() { + EmbeddedChannel channel = newChannel(); + ChannelHandlerContext ctx = context(channel); - when(ctx.channel()).thenReturn(channel); - when(msg.variableHeader()).thenReturn(variableHeader); - when(msg.payload()).thenReturn(payload); - when(variableHeader.version()).thenReturn((int) MqttVersion.MQTT_3_1.protocolLevel()); - when(payload.clientIdentifier()).thenReturn("test-client-001"); - when(payload.userName()).thenReturn(VALID_USER); - when(payload.passwordInBytes()).thenReturn(VALID_PASS.getBytes()); + connect.connect(ctx, connectMessage(MqttVersion.MQTT_3_1_1)); + assertNotNull(channel.readOutbound()); - MqttContext mqttContext = new MqttContext(); - mqttContext.setUserName(VALID_USER); - mqttContext.setPassword(VALID_PASS); - } + connect.connect(ctx, connectMessage(MqttVersion.MQTT_3_1_1)); - @AfterEach - public void tearDown() { - Singleton.INST.single(WillRepository.class, new WillRepository()); - Singleton.INST.single(ChannelRepository.class, new ChannelRepository()); + assertChannelClosed(channel); + assertNull(channel.readOutbound()); } @Test public void testStoresWillOnConnect() { - byte[] willMessage = "client disconnected unexpectedly".getBytes(); - when(variableHeader.isWillFlag()).thenReturn(true); - when(variableHeader.willQos()).thenReturn(1); - when(variableHeader.isWillRetain()).thenReturn(true); - when(payload.willTopic()).thenReturn("status/client-001"); - when(payload.willMessageInBytes()).thenReturn(willMessage); - - connect.connect(ctx, msg); - - WillRepository.WillEntry will = willRepository.get(channel); - assertThat(will, notNullValue()); - assertEquals("status/client-001", will.getTopic()); + 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()); @@ -226,45 +197,98 @@ public void testStoresWillOnConnect() { @Test public void testDoesNotStoreWillWhenWillFlagIsFalse() { - when(variableHeader.isWillFlag()).thenReturn(false); + EmbeddedChannel channel = newChannel(); - connect.connect(ctx, msg); + connect.connect(context(channel), connectMessage(MqttVersion.MQTT_3_1_1)); - WillRepository.WillEntry will = willRepository.get(channel); - assertThat(will, nullValue()); + assertNull(willRepository.get(channel)); } @Test public void testWillQosZero() { - byte[] willMessage = "qos0 will".getBytes(); - when(variableHeader.isWillFlag()).thenReturn(true); - when(variableHeader.willQos()).thenReturn(0); - when(variableHeader.isWillRetain()).thenReturn(false); - when(payload.willTopic()).thenReturn("topic/qos0"); - when(payload.willMessageInBytes()).thenReturn(willMessage); + EmbeddedChannel channel = newChannel(); + byte[] willMessage = "qos0 will".getBytes(StandardCharsets.UTF_8); - connect.connect(ctx, msg); + connect.connect(context(channel), willConnectMessage(0, false, "topic/qos0", willMessage)); - WillRepository.WillEntry will = willRepository.get(channel); - assertThat(will, notNullValue()); + WillEntry will = willRepository.get(channel); + assertNotNull(will); assertEquals(0, will.getQos()); assertFalse(will.isRetain()); } @Test public void testWillRetainTrue() { - byte[] willMessage = "retained will".getBytes(); - when(variableHeader.isWillFlag()).thenReturn(true); - when(variableHeader.willQos()).thenReturn(2); - when(variableHeader.isWillRetain()).thenReturn(true); - when(payload.willTopic()).thenReturn("topic/retained"); - when(payload.willMessageInBytes()).thenReturn(willMessage); + EmbeddedChannel channel = newChannel(); + byte[] willMessage = "retained will".getBytes(StandardCharsets.UTF_8); - connect.connect(ctx, msg); + connect.connect(context(channel), willConnectMessage(2, true, "topic/retained", willMessage)); - WillRepository.WillEntry will = willRepository.get(channel); - assertThat(will, notNullValue()); - assertTrue(will.isRetain()); + 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 = newChannel(); + + connect.connect(context(channel), connectMessage(version)); + + MqttConnAckMessage ackMessage = channel.readOutbound(); + assertNotNull(ackMessage); + assertEquals(CONNECTION_ACCEPTED, ackMessage.variableHeader().connectReturnCode()); + assertTrue(ackMessage.variableHeader().isSessionPresent()); + 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, 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/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 3086f3840467..0ac538b2be39 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 @@ -27,6 +27,7 @@ import io.netty.handler.codec.mqtt.MqttVersion; import io.netty.channel.Channel; import io.netty.channel.ChannelHandlerContext; +import io.netty.channel.ChannelInboundHandlerAdapter; import org.apache.shenyu.common.utils.Singleton; import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository; import org.junit.jupiter.api.AfterAll; @@ -41,8 +42,10 @@ import org.mockito.junit.jupiter.MockitoExtension; import java.nio.charset.StandardCharsets; +import java.util.concurrent.atomic.AtomicInteger; import static org.hamcrest.MatcherAssert.assertThat; import static org.hamcrest.Matchers.nullValue; +import static org.mockito.Mockito.lenient; import static org.mockito.Mockito.when; import static org.junit.jupiter.api.Assertions.assertEquals; @@ -77,7 +80,7 @@ public final class MqttTransportHandlerTest { private WillRepository willRepository; @BeforeAll - static void setUp() { + static void setUpAll() { channelRepository = new ChannelRepository(); Singleton.INST.single(ChannelRepository.class, channelRepository); new MqttContext().setUserName(USER_NAME); @@ -85,7 +88,7 @@ static void setUp() { } @AfterAll - static void tearDown() { + static void tearDownAll() { new MqttContext().setUserName(null); new MqttContext().setPassword(null); } @@ -120,6 +123,24 @@ public void abruptChannelCloseCleansUpChannelRepository() { channel.finishAndReleaseAll(); } + @Test + public void channelInactiveIsPropagatedOnlyOnce() { + AtomicInteger fired = new AtomicInteger(); + EmbeddedChannel channel = new EmbeddedChannel(new MqttTransportHandler(), new ChannelInboundHandlerAdapter() { + @Override + public void channelInactive(final ChannelHandlerContext context) throws Exception { + fired.incrementAndGet(); + super.channelInactive(context); + } + }); + + channel.close(); + channel.runPendingTasks(); + + assertEquals(1, fired.get()); + channel.finishAndReleaseAll(); + } + private MqttConnectMessage connectMessage() { MqttFixedHeader fixedHeader = new MqttFixedHeader(MqttMessageType.CONNECT, false, MqttQoS.AT_MOST_ONCE, false, 0); MqttConnectVariableHeader variableHeader = new MqttConnectVariableHeader( @@ -131,16 +152,17 @@ private MqttConnectMessage connectMessage() { } @BeforeEach - public void setUp2() { + public void setUp() { handler = new MqttTransportHandler(); willRepository = new WillRepository(); Singleton.INST.single(WillRepository.class, willRepository); Singleton.INST.single(SubscribeRepository.class, subscribeRepository); - when(ctx.channel()).thenReturn(channel); + // shared stub: only the tests driving the handler with a mock context use it + lenient().when(ctx.channel()).thenReturn(channel); } @AfterEach - public void tearDown2() { + public void tearDown() { Singleton.INST.single(WillRepository.class, new WillRepository()); Singleton.INST.single(SubscribeRepository.class, new SubscribeRepository()); } 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 3a629cc7699f..2320ce522764 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 @@ -26,20 +26,17 @@ import java.util.ArrayList; import java.util.Arrays; +import java.util.Collections; +import java.util.LinkedHashSet; import java.util.List; -import java.util.function.BooleanSupplier; +import java.util.Set; +import java.util.concurrent.CountDownLatch; import java.util.stream.Collectors; -import java.time.Duration; -import java.util.Collections; -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.junit.jupiter.api.Assertions.fail; -import static org.mockito.Mockito.mock; /** * Test cases for {@link SubscribeRepository}. @@ -52,52 +49,45 @@ public final class SubscribeRepositoryTest { private static final String KEPT_TOPIC = "test/kept-topic"; + private static final long JOIN_TIMEOUT_MILLIS = 5000L; + private final SubscribeRepository repository = new SubscribeRepository(); private final List channels = new ArrayList<>(); - private final List topicFilters = new ArrayList<>(); + private final Set subscribedTopics = new LinkedHashSet<>(); + + @AfterEach + public void tearDown() { + // subscriptions live in a static map shared with the other test classes, + // so every topic has to be empty before the next test runs. + channels.forEach(channel -> subscribedTopics.forEach( + topic -> repository.remove(Collections.singletonList(topic), channel))); + subscribedTopics.forEach(topic -> assertTrue(repository.get(topic).isEmpty())); + channels.forEach(channel -> ((EmbeddedChannel) channel).finishAndReleaseAll()); + } @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)); + Channel channel = newSubscriber(EXISTING_TOPIC); repository.remove(Collections.singletonList(EXISTING_TOPIC), channel); - await().atMost(Duration.ofSeconds(5)).until(() -> repository.get(EXISTING_TOPIC).isEmpty()); + assertTrue(repository.get(EXISTING_TOPIC).isEmpty()); } @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)); + Channel channel = newSubscriber(KEPT_TOPIC); assertDoesNotThrow(() -> repository.remove(Collections.singletonList(ABSENT_TOPIC), channel)); - await().atMost(Duration.ofSeconds(5)) - .until(() -> ForkJoinPool.commonPool().awaitQuiescence(1, TimeUnit.SECONDS)); assertTrue(repository.get(ABSENT_TOPIC).isEmpty()); assertTrue(repository.get(KEPT_TOPIC).contains(channel)); } - @AfterEach - void cleanup() throws InterruptedException { - if (!topicFilters.isEmpty()) { - for (Channel subscribed : channels) { - repository.remove(topicFilters, subscribed); - } - awaitUntil(() -> topicFilters.stream().allMatch(topic -> repository.get(topic).isEmpty())); - } - channels.forEach(channel -> ((EmbeddedChannel) channel).finishAndReleaseAll()); - } - @Test - void testGetChannelsByTopicExactMatch() throws InterruptedException { + public void testGetChannelsByTopicExactMatch() { Channel subscriber = newSubscriber("sport/tennis"); assertTrue(repository.getChannelsByTopic("sport/tennis").contains(subscriber)); @@ -105,7 +95,7 @@ void testGetChannelsByTopicExactMatch() throws InterruptedException { } @Test - void testGetChannelsByTopicWildcardMatch() throws InterruptedException { + public void testGetChannelsByTopicWildcardMatch() { Channel subscriber = newSubscriber("sport/+/player1"); assertTrue(repository.getChannelsByTopic("sport/tennis/player1").contains(subscriber)); @@ -113,7 +103,7 @@ void testGetChannelsByTopicWildcardMatch() throws InterruptedException { } @Test - void testGetChannelsByTopicDeduplicatesOverlappingSubscriptions() throws InterruptedException { + public void testGetChannelsByTopicDeduplicatesOverlappingSubscriptions() { Channel subscriber = newSubscriber("#", "sport/#"); // MQTT requires at most one delivery per publish per client @@ -122,23 +112,48 @@ void testGetChannelsByTopicDeduplicatesOverlappingSubscriptions() throws Interru } @Test - void testGetChannelsByTopicMultipleSubscribers() throws InterruptedException { - final Channel first = newSubscriber("sport/tennis"); - final Channel second = new EmbeddedChannel(); - channels.add(second); - topicFilters.add("sport/tennis"); - repository.add(second, subscriptions("sport/tennis")); - awaitUntil(() -> repository.get("sport/tennis").containsAll(Arrays.asList(first, second))); + public void testGetChannelsByTopicMultipleSubscribers() { + Channel first = newSubscriber("sport/tennis"); + Channel second = newSubscriber("sport/tennis"); assertEquals(2, repository.getChannelsByTopic("sport/tennis").size()); + assertTrue(repository.getChannelsByTopic("sport/tennis").containsAll(Arrays.asList(first, second))); + } + + /** + * Regression guard: subscribing to a brand new topic must keep every client + * when several clients send SUBSCRIBE at the same time. + */ + @Test + public void concurrentSubscribersOfTheSameNewTopicAreAllRegistered() throws InterruptedException { + int subscriberCount = 8; + String topic = "sport/concurrent"; + subscribedTopics.add(topic); + CountDownLatch startGate = new CountDownLatch(1); + List subscribingThreads = new ArrayList<>(); + for (int i = 0; i < subscriberCount; i++) { + EmbeddedChannel subscriber = new EmbeddedChannel(); + channels.add(subscriber); + Thread thread = new Thread(() -> subscribeAfter(startGate, subscriber, 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"); + } + + assertEquals(subscriberCount, repository.get(topic).size()); + assertEquals(subscriberCount, repository.getChannelsByTopic(topic).size()); } - private Channel newSubscriber(final String... topics) throws InterruptedException { + private Channel newSubscriber(final String... topics) { Channel subscriber = new EmbeddedChannel(); channels.add(subscriber); - topicFilters.addAll(Arrays.asList(topics)); + subscribedTopics.addAll(Arrays.asList(topics)); repository.add(subscriber, subscriptions(topics)); - awaitUntil(() -> Arrays.stream(topics).allMatch(topic -> repository.get(topic).contains(subscriber))); return subscriber; } @@ -148,13 +163,13 @@ private List subscriptions(final String... topics) { .collect(Collectors.toList()); } - private void awaitUntil(final BooleanSupplier condition) throws InterruptedException { - long deadline = System.currentTimeMillis() + 5000; - while (!condition.getAsBoolean()) { - if (System.currentTimeMillis() >= deadline) { - fail("condition not met within timeout"); - } - Thread.sleep(10); + 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, subscriptions(topic)); } } From b29821d3b164803f367335e91bbb2fb5d447fa77 Mon Sep 17 00:00:00 2001 From: wy471x Date: Fri, 2 Oct 2026 01:32:32 +0800 Subject: [PATCH 5/6] fix(mqtt): resolve merge conflicts in SubscribeRepository and its tests - adapt getChannelsByTopic to the QoS-aware Map storage - merge both sides of MqttTransportHandlerTest into a single lifecycle - merge SubscribeRepositoryTest onto the async QoS API and drop duplicates - remove duplicated imports and unused imports left by the merge --- .../repositories/SubscribeRepository.java | 16 +- .../mqtt/MqttTransportHandlerTest.java | 203 ++++++++---------- .../repositories/SubscribeRepositoryTest.java | 187 +++++++--------- 3 files changed, 169 insertions(+), 237 deletions(-) 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 86e4b2381c3e..b0e09121a7b3 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,23 +20,19 @@ 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; -import org.apache.commons.collections4.CollectionUtils; -import org.apache.shenyu.protocol.mqtt.TopicMatcher; -import java.util.Collections; import java.util.ArrayList; import java.util.Collections; import java.util.LinkedHashSet; import java.util.List; import java.util.Map; import java.util.Objects; -import java.util.concurrent.CompletableFuture; -import java.util.Objects; import java.util.Set; +import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; -import java.util.concurrent.CopyOnWriteArrayList; /** * Topic and channel association. @@ -130,18 +126,18 @@ public List getChannelsByTopic(final String topic) { Set result = new LinkedHashSet<>(); // fast path: exact subscription, no wildcard scan needed - List exactMatch = TOPIC_CHANNEL_FACTORY.get(topic); + Map exactMatch = TOPIC_CHANNEL_FACTORY.get(topic); if (Objects.nonNull(exactMatch)) { - result.addAll(exactMatch); + result.addAll(exactMatch.keySet()); } - for (Map.Entry> entry : TOPIC_CHANNEL_FACTORY.entrySet()) { + 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)) { - result.addAll(entry.getValue()); + result.addAll(entry.getValue().keySet()); } } return new ArrayList<>(result); 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 522211933a96..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; @@ -30,25 +33,18 @@ import io.netty.handler.codec.mqtt.MqttQoS; import io.netty.handler.codec.mqtt.MqttTopicSubscription; import io.netty.handler.codec.mqtt.MqttVersion; -import io.netty.channel.Channel; -import io.netty.channel.ChannelHandlerContext; -import io.netty.channel.ChannelInboundHandlerAdapter; import io.netty.util.CharsetUtil; 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.repositories.WillRepository; -import org.junit.jupiter.api.AfterEach; -import org.junit.jupiter.api.BeforeEach; -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 org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; @@ -56,16 +52,15 @@ import java.time.Duration; import java.util.Collections; import java.util.concurrent.atomic.AtomicInteger; -import static org.hamcrest.MatcherAssert.assertThat; -import static org.hamcrest.Matchers.nullValue; -import static org.mockito.Mockito.lenient; -import static org.mockito.Mockito.when; 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}. @@ -75,14 +70,14 @@ 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"; 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); @@ -95,27 +90,29 @@ public final class MqttTransportHandlerTest { private static final SubscribeRepository SUBSCRIBE_REPOSITORY = new SubscribeRepository(); - private EmbeddedChannel registeredChannel; - @Mock private ChannelHandlerContext ctx; @Mock private Channel channel; - @Mock - private SubscribeRepository subscribeRepository; - 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); @@ -129,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 @@ -210,44 +212,10 @@ public void testOperationCompleteCleansRepositoriesOnClose() throws Exception { assertEquals(1, MqttPacketIdGenerator.next(registeredChannel)); } - private MqttConnectMessage connectMessage() { - MqttFixedHeader fixedHeader = new MqttFixedHeader(MqttMessageType.CONNECT, false, MqttQoS.AT_MOST_ONCE, false, 0); - MqttConnectVariableHeader variableHeader = new MqttConnectVariableHeader( - MqttVersion.MQTT_3_1_1.protocolName(), MqttVersion.MQTT_3_1_1.protocolLevel(), - true, true, false, 0, false, false, 60); - MqttConnectPayload payload = new MqttConnectPayload(CLIENT_ID, null, null, - USER_NAME, PASSWORD.getBytes(StandardCharsets.UTF_8)); - return new MqttConnectMessage(fixedHeader, variableHeader, payload); - } - - /** - * 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); - } - - @BeforeAll - static void setUpAll() { - channelRepository = new ChannelRepository(); - Singleton.INST.single(ChannelRepository.class, channelRepository); - new MqttContext().setUserName(USER_NAME); - new MqttContext().setPassword(PASSWORD); - } - - @AfterAll - static void tearDownAll() { - new MqttContext().setUserName(null); - new MqttContext().setPassword(null); - } - @Test public void channelInactiveIsPropagatedOnlyOnce() { AtomicInteger fired = new AtomicInteger(); - EmbeddedChannel channel = new EmbeddedChannel(new MqttTransportHandler(), new ChannelInboundHandlerAdapter() { + EmbeddedChannel sessionChannel = new EmbeddedChannel(new MqttTransportHandler(), new ChannelInboundHandlerAdapter() { @Override public void channelInactive(final ChannelHandlerContext context) throws Exception { fired.incrementAndGet(); @@ -255,61 +223,66 @@ public void channelInactive(final ChannelHandlerContext context) throws Exceptio } }); - channel.close(); - channel.runPendingTasks(); + sessionChannel.close(); + sessionChannel.runPendingTasks(); assertEquals(1, fired.get()); - channel.finishAndReleaseAll(); - } - - @BeforeEach - public void setUpEach() { - handler = new MqttTransportHandler(); - willRepository = new WillRepository(); - Singleton.INST.single(WillRepository.class, willRepository); - Singleton.INST.single(SubscribeRepository.class, subscribeRepository); - // shared stub: only the tests driving the handler with a mock context use it - lenient().when(ctx.channel()).thenReturn(channel); - } - - @AfterEach - public void tearDownEach() { - Singleton.INST.single(WillRepository.class, new WillRepository()); - Singleton.INST.single(SubscribeRepository.class, new SubscribeRepository()); + sessionChannel.finishAndReleaseAll(); } @Test public void testChannelInactiveFiresWillAndRemovesIt() throws Exception { - byte[] willMessage = "sudden disconnect".getBytes(); - WillRepository.WillEntry will = new WillRepository.WillEntry("status/offline", willMessage, 1, true); - willRepository.add(channel, will); - - // publishWill uses subscribeRepository to get target channels - when(subscribeRepository.getChannelsByTopic("status/offline")).thenReturn(java.util.Collections.emptyList()); + 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); - // will should be removed after firing - assertThat(willRepository.get(channel), nullValue()); + 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); - assertThat(willRepository.get(channel), nullValue()); + assertNull(willRepository.get(channel)); } @Test public void testChannelInactiveAfterDisconnectClearsWill() throws Exception { - byte[] willMessage = "graceful close".getBytes(); - WillRepository.WillEntry will = new WillRepository.WillEntry("status/clean", willMessage, 0, false); - willRepository.add(channel, will); + willRepository.add(channel, new WillRepository.WillEntry("status/clean", "graceful close".getBytes(), 0, false)); - // simulate graceful disconnect: remove will first, then channelInactive + // a graceful disconnect removes the will first, so channelInactive must not publish it willRepository.remove(channel); handler.channelInactive(ctx); - assertThat(willRepository.get(channel), nullValue()); + assertNull(willRepository.get(channel)); + } + + private MqttConnectMessage connectMessage() { + MqttFixedHeader fixedHeader = new MqttFixedHeader(MqttMessageType.CONNECT, false, MqttQoS.AT_MOST_ONCE, false, 0); + MqttConnectVariableHeader variableHeader = new MqttConnectVariableHeader( + MqttVersion.MQTT_3_1_1.protocolName(), MqttVersion.MQTT_3_1_1.protocolLevel(), + true, true, false, 0, false, false, 60); + MqttConnectPayload payload = new MqttConnectPayload(CLIENT_ID, null, null, + USER_NAME, PASSWORD.getBytes(StandardCharsets.UTF_8)); + return new MqttConnectMessage(fixedHeader, variableHeader, payload); + } + + /** + * 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/repositories/SubscribeRepositoryTest.java b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepositoryTest.java index e89963fd943e..995a21cba155 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 @@ -18,31 +18,24 @@ package org.apache.shenyu.protocol.mqtt.repositories; import io.netty.channel.Channel; -import io.netty.channel.embedded.EmbeddedChannel; 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.AfterEach; import org.junit.jupiter.api.Test; import java.time.Duration; -import java.util.Arrays; 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; -import java.util.LinkedHashSet; -import java.util.List; -import java.util.Set; -import java.util.concurrent.CountDownLatch; -import java.util.stream.Collectors; import static org.awaitility.Awaitility.await; import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; @@ -56,17 +49,27 @@ */ 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 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); @@ -74,11 +77,7 @@ public final class SubscribeRepositoryTest { private static final long JOIN_TIMEOUT_MILLIS = 5000L; - private final SubscribeRepository repository = new SubscribeRepository(); - - private final List channels = new ArrayList<>(); - - private final Set subscribedTopics = new LinkedHashSet<>(); + private SubscribeRepository repository; private Channel channel; @@ -216,111 +215,60 @@ public void testRemoveAbsentTopicDoesNotThrow() { 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 not 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()))); - } - - @AfterEach - public void tearDown() { - // subscriptions live in a static map shared with the other test classes, - // so every topic has to be empty before the next test runs. - channels.forEach(channel -> subscribedTopics.forEach( - topic -> repository.remove(Collections.singletonList(topic), channel))); - subscribedTopics.forEach(topic -> assertTrue(repository.get(topic).isEmpty())); - channels.forEach(channel -> ((EmbeddedChannel) channel).finishAndReleaseAll()); - } - - @Test - public void removeRemovesChannelFromExistingTopic() { - Channel channel = newSubscriber(EXISTING_TOPIC); - - repository.remove(Collections.singletonList(EXISTING_TOPIC), channel); - - assertTrue(repository.get(EXISTING_TOPIC).isEmpty()); - } - - @Test - public void removeAbsentTopicDoesNotThrow() { - Channel channel = newSubscriber(KEPT_TOPIC); - - assertDoesNotThrow(() -> repository.remove(Collections.singletonList(ABSENT_TOPIC), channel)); - - assertTrue(repository.get(ABSENT_TOPIC).isEmpty()); - assertTrue(repository.get(KEPT_TOPIC).contains(channel)); - } - @Test public void testGetChannelsByTopicExactMatch() { - Channel subscriber = newSubscriber("sport/tennis"); - - assertTrue(repository.getChannelsByTopic("sport/tennis").contains(subscriber)); - assertTrue(repository.getChannelsByTopic("sport/tennis/player1").isEmpty()); + repository.add(channel, Collections.singletonList(new MqttTopicSubscription(EXACT_TOPIC, MqttQoS.AT_MOST_ONCE))); + awaitAssert(() -> assertTrue(repository.getChannelsByTopic(EXACT_TOPIC).contains(channel))); + assertTrue(repository.getChannelsByTopic(CHILD_TOPIC).isEmpty()); } @Test public void testGetChannelsByTopicWildcardMatch() { - Channel subscriber = newSubscriber("sport/+/player1"); - - assertTrue(repository.getChannelsByTopic("sport/tennis/player1").contains(subscriber)); - assertTrue(repository.getChannelsByTopic("sport/tennis").isEmpty()); + repository.add(channel, Collections.singletonList(new MqttTopicSubscription(SINGLE_LEVEL_FILTER, MqttQoS.AT_MOST_ONCE))); + awaitAssert(() -> assertTrue(repository.getChannelsByTopic(CHILD_TOPIC).contains(channel))); + assertTrue(repository.getChannelsByTopic(EXACT_TOPIC).isEmpty()); } @Test public void testGetChannelsByTopicDeduplicatesOverlappingSubscriptions() { - Channel subscriber = newSubscriber("#", "sport/#"); + repository.add(channel, Arrays.asList( + new MqttTopicSubscription(MATCH_ALL_FILTER, MqttQoS.AT_MOST_ONCE), + new MqttTopicSubscription(MULTI_LEVEL_FILTER, MqttQoS.AT_MOST_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 - assertEquals(1, repository.getChannelsByTopic("sport/tennis").size()); - assertTrue(repository.getChannelsByTopic("sport/tennis").contains(subscriber)); + List matched = repository.getChannelsByTopic(EXACT_TOPIC); + assertEquals(1, matched.size()); + assertTrue(matched.contains(channel)); } @Test public void testGetChannelsByTopicMultipleSubscribers() { - Channel first = newSubscriber("sport/tennis"); - Channel second = newSubscriber("sport/tennis"); - - assertEquals(2, repository.getChannelsByTopic("sport/tennis").size()); - assertTrue(repository.getChannelsByTopic("sport/tennis").containsAll(Arrays.asList(first, second))); + 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).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; - String topic = "sport/concurrent"; - subscribedTopics.add(topic); CountDownLatch startGate = new CountDownLatch(1); + List subscribers = new ArrayList<>(); List subscribingThreads = new ArrayList<>(); for (int i = 0; i < subscriberCount; i++) { - EmbeddedChannel subscriber = new EmbeddedChannel(); - channels.add(subscriber); - Thread thread = new Thread(() -> subscribeAfter(startGate, subscriber, topic)); + Channel subscriber = mock(Channel.class); + subscribers.add(subscriber); + Thread thread = new Thread(() -> subscribeAfter(startGate, subscriber, CONCURRENT_TOPIC)); thread.start(); subscribingThreads.add(thread); } @@ -331,22 +279,10 @@ public void concurrentSubscribersOfTheSameNewTopicAreAllRegistered() throws Inte assertFalse(thread.isAlive(), "subscribing thread did not finish in time"); } - assertEquals(subscriberCount, repository.get(topic).size()); - assertEquals(subscriberCount, repository.getChannelsByTopic(topic).size()); - } - - private Channel newSubscriber(final String... topics) { - Channel subscriber = new EmbeddedChannel(); - channels.add(subscriber); - subscribedTopics.addAll(Arrays.asList(topics)); - repository.add(subscriber, subscriptions(topics)); - return subscriber; - } - - private List subscriptions(final String... topics) { - return Arrays.stream(topics) - .map(topic -> new MqttTopicSubscription(topic, MqttQoS.AT_MOST_ONCE)) - .collect(Collectors.toList()); + awaitAssert(() -> assertEquals(subscriberCount, repository.get(CONCURRENT_TOPIC).size())); + List matched = repository.getChannelsByTopic(CONCURRENT_TOPIC); + assertEquals(subscriberCount, matched.size()); + assertTrue(matched.containsAll(subscribers)); } private void subscribeAfter(final CountDownLatch startGate, final Channel subscriber, final String topic) { @@ -356,6 +292,33 @@ private void subscribeAfter(final CountDownLatch startGate, final Channel subscr Thread.currentThread().interrupt(); return; } - repository.add(subscriber, subscriptions(topic)); + 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. + * + * @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 not 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()))); } } From 71fb63a474b00b0b1b837f347792866476721342 Mon Sep 17 00:00:00 2001 From: wy471x Date: Fri, 2 Oct 2026 07:35:39 +0800 Subject: [PATCH 6/6] fix(mqtt): cap will qos at the subscriber granted qos getChannelsByTopic now returns matching channels mapped to the maximum granted qos across exact, wildcard and overlapping filters, and publishWill delivers each will with min(will qos, granted qos) like the normal publish path, so a qos 2 will is no longer sent at qos 2 to a qos 0 subscriber. --- .../apache/shenyu/protocol/mqtt/Publish.java | 15 +++---- .../repositories/SubscribeRepository.java | 22 +++++----- .../shenyu/protocol/mqtt/PublishWillTest.java | 41 ++++++++++++++++--- .../repositories/SubscribeRepositoryTest.java | 18 ++++---- 4 files changed, 63 insertions(+), 33 deletions(-) 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 277e7e1ca11b..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 @@ -37,7 +37,6 @@ import org.apache.shenyu.protocol.mqtt.repositories.WillRepository; -import java.util.List; import java.util.Objects; import java.util.Map; import java.util.concurrent.CompletableFuture; @@ -157,14 +156,16 @@ static void publishWill(final WillRepository.WillEntry will) { if (Objects.isNull(will) || Objects.isNull(will.getTopic()) || Objects.isNull(will.getMessage())) { return; } - final List channels = Singleton.INST.get(SubscribeRepository.class).getChannelsByTopic(will.getTopic()); + final Map subscribers = Singleton.INST.get(SubscribeRepository.class).getChannelsByTopic(will.getTopic()); final MqttQoS willQos = MqttQoS.valueOf(will.getQos()); - final int packetId = willQos == MqttQoS.AT_MOST_ONCE - ? 0 - : java.util.concurrent.ThreadLocalRandom.current().nextInt(1, 65536); - channels.parallelStream().forEach(channel -> { + subscribers.entrySet().parallelStream().forEach(entry -> { + Channel channel = entry.getKey(); if (channel.isActive()) { - MqttFixedHeader mqttFixedHeader = new MqttFixedHeader(MqttMessageType.PUBLISH, false, willQos, will.isRetain(), 0); + 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())); 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 b0e09121a7b3..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 @@ -24,13 +24,10 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import java.util.ArrayList; import java.util.Collections; -import java.util.LinkedHashSet; import java.util.List; import java.util.Map; import java.util.Objects; -import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; @@ -114,21 +111,22 @@ private static MqttQoS maxQoS(final MqttQoS qos1, final MqttQoS qos2) { } /** - * Get channels whose subscription filter matches the published topic. + * 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 channels subscribed to matching topic filters + * @return matching channels with their maximum granted qos */ - public List getChannelsByTopic(final String topic) { - // MQTT requires at most one delivery per publish per client, so dedupe - // channels when overlapping filters (e.g. sport/# and #) both match. - Set result = new LinkedHashSet<>(); + 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.addAll(exactMatch.keySet()); + result.putAll(exactMatch); } for (Map.Entry> entry : TOPIC_CHANNEL_FACTORY.entrySet()) { @@ -137,10 +135,10 @@ public List getChannelsByTopic(final String topic) { continue; } if (TopicMatcher.matches(filter, topic)) { - result.addAll(entry.getValue().keySet()); + entry.getValue().forEach((channel, qos) -> result.merge(channel, qos, SubscribeRepository::maxQoS)); } } - return new ArrayList<>(result); + return result; } } 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 index 7ba6d1590c7d..4e89d1e8e63d 100644 --- 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 @@ -19,6 +19,7 @@ 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; @@ -31,6 +32,8 @@ 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; @@ -49,6 +52,9 @@ public class PublishWillTest { @Mock private Channel subscriberChannel; + @Mock + private Channel otherChannel; + @BeforeEach public void setUp() { Singleton.INST.single(SubscribeRepository.class, subscribeRepository); @@ -63,7 +69,7 @@ public void tearDown() { public void testPublishWillToActiveSubscriber() { when(subscriberChannel.isActive()).thenReturn(true); when(subscribeRepository.getChannelsByTopic("status/offline")) - .thenReturn(Collections.singletonList(subscriberChannel)); + .thenReturn(Collections.singletonMap(subscriberChannel, MqttQoS.EXACTLY_ONCE)); byte[] message = "client lost".getBytes(); WillRepository.WillEntry will = new WillRepository.WillEntry("status/offline", message, 1, true); @@ -81,7 +87,7 @@ public void testPublishWillToActiveSubscriber() { public void testPublishWillSkipsInactiveChannel() { when(subscriberChannel.isActive()).thenReturn(false); when(subscribeRepository.getChannelsByTopic("status/inactive")) - .thenReturn(Collections.singletonList(subscriberChannel)); + .thenReturn(Collections.singletonMap(subscriberChannel, MqttQoS.AT_LEAST_ONCE)); WillRepository.WillEntry will = new WillRepository.WillEntry("status/inactive", "msg".getBytes(), 0, false); Publish.publishWill(will); @@ -91,7 +97,7 @@ public void testPublishWillSkipsInactiveChannel() { @Test public void testPublishWillToEmptySubscribers() { - when(subscribeRepository.getChannelsByTopic("topic/none")).thenReturn(Collections.emptyList()); + when(subscribeRepository.getChannelsByTopic("topic/none")).thenReturn(Collections.emptyMap()); WillRepository.WillEntry will = new WillRepository.WillEntry("topic/none", "msg".getBytes(), 2, false); Publish.publishWill(will); @@ -101,7 +107,7 @@ public void testPublishWillToEmptySubscribers() { public void testPublishWillQosAndRetain() { when(subscriberChannel.isActive()).thenReturn(true); when(subscribeRepository.getChannelsByTopic("qos/retain")) - .thenReturn(Collections.singletonList(subscriberChannel)); + .thenReturn(Collections.singletonMap(subscriberChannel, MqttQoS.EXACTLY_ONCE)); WillRepository.WillEntry will = new WillRepository.WillEntry("qos/retain", "data".getBytes(), 0, false); Publish.publishWill(will); @@ -117,7 +123,7 @@ public void testPublishWillQosAndRetain() { public void testPublishWillToWildcardSubscriber() { when(subscriberChannel.isActive()).thenReturn(true); when(subscribeRepository.getChannelsByTopic("status/client-001")) - .thenReturn(Collections.singletonList(subscriberChannel)); + .thenReturn(Collections.singletonMap(subscriberChannel, MqttQoS.AT_MOST_ONCE)); WillRepository.WillEntry will = new WillRepository.WillEntry("status/client-001", "gone".getBytes(), 0, false); Publish.publishWill(will); @@ -127,4 +133,29 @@ public void testPublishWillToWildcardSubscriber() { 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/repositories/SubscribeRepositoryTest.java b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/repositories/SubscribeRepositoryTest.java index 995a21cba155..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 @@ -218,14 +218,14 @@ public void testRemoveAbsentTopicDoesNotThrow() { @Test public void testGetChannelsByTopicExactMatch() { repository.add(channel, Collections.singletonList(new MqttTopicSubscription(EXACT_TOPIC, MqttQoS.AT_MOST_ONCE))); - awaitAssert(() -> assertTrue(repository.getChannelsByTopic(EXACT_TOPIC).contains(channel))); + 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).contains(channel))); + awaitAssert(() -> assertTrue(repository.getChannelsByTopic(CHILD_TOPIC).containsKey(channel))); assertTrue(repository.getChannelsByTopic(EXACT_TOPIC).isEmpty()); } @@ -233,16 +233,16 @@ public void testGetChannelsByTopicWildcardMatch() { public void testGetChannelsByTopicDeduplicatesOverlappingSubscriptions() { repository.add(channel, Arrays.asList( new MqttTopicSubscription(MATCH_ALL_FILTER, MqttQoS.AT_MOST_ONCE), - new MqttTopicSubscription(MULTI_LEVEL_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 - List matched = repository.getChannelsByTopic(EXACT_TOPIC); + // 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()); - assertTrue(matched.contains(channel)); + assertEquals(MqttQoS.EXACTLY_ONCE, matched.get(channel)); } @Test @@ -250,7 +250,7 @@ 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).containsAll(Arrays.asList(channel, otherChannel))); + assertTrue(repository.getChannelsByTopic(EXACT_TOPIC).keySet().containsAll(Arrays.asList(channel, otherChannel))); } /** @@ -280,9 +280,9 @@ public void concurrentSubscribersOfTheSameNewTopicAreAllRegistered() throws Inte } awaitAssert(() -> assertEquals(subscriberCount, repository.get(CONCURRENT_TOPIC).size())); - List matched = repository.getChannelsByTopic(CONCURRENT_TOPIC); + Map matched = repository.getChannelsByTopic(CONCURRENT_TOPIC); assertEquals(subscriberCount, matched.size()); - assertTrue(matched.containsAll(subscribers)); + assertTrue(matched.keySet().containsAll(subscribers)); } private void subscribeAfter(final CountDownLatch startGate, final Channel subscriber, final String topic) {