This is an automated email from the ASF dual-hosted git repository.

Aias00 pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/shenyu.git


The following commit(s) were added to refs/heads/master by this push:
     new b0b7fab03b fix: enforce per-channel MQTT connection state guard (#6983)
b0b7fab03b is described below

commit b0b7fab03b198c561c6f9f363755ff8eae9b1913
Author: wy471x <[email protected]>
AuthorDate: Fri Sep 4 11:20:53 2026 +0800

    fix: enforce per-channel MQTT connection state guard (#6983)
    
    * fix: enforce mqtt connection state guard per channel to reject 
pre-connect operations and duplicate connect
    
    Co-Authored-By: Claude <[email protected]>
    
    * fix: clean up channel repository entry when mqtt channel becomes inactive
    
    A channel registered by a successful CONNECT stayed in ChannelRepository
    when it was closed outside the explicit DISCONNECT flow, e.g. after a
    duplicate CONNECT rejected by the connection state guard or an abrupt
    client disconnect. Clean up centrally in channelInactive, and make the
    registration synchronous so an in-flight put can not re-add a closed
    channel after the removal.
    
    ---------
    
    Co-authored-by: Claude <[email protected]>
    Co-authored-by: aias00 <[email protected]>
---
 .../org/apache/shenyu/protocol/mqtt/Connect.java   |  8 ++-
 .../apache/shenyu/protocol/mqtt/MessageType.java   | 14 ++--
 .../shenyu/protocol/mqtt/MqttTransportHandler.java |  8 +++
 .../org/apache/shenyu/protocol/mqtt/PingReq.java   |  6 ++
 .../org/apache/shenyu/protocol/mqtt/Publish.java   |  4 +-
 .../org/apache/shenyu/protocol/mqtt/Subscribe.java |  2 +-
 .../apache/shenyu/protocol/mqtt/Unsubscribe.java   |  4 +-
 .../mqtt/repositories/ChannelRepository.java       |  3 +-
 .../apache/shenyu/protocol/mqtt/ConnectTest.java   | 54 +++++++++------
 ...nectTest.java => MqttTransportHandlerTest.java} | 81 ++++++++--------------
 .../apache/shenyu/protocol/mqtt/PingReqTest.java   | 56 +++++++++++++++
 .../apache/shenyu/protocol/mqtt/PublishTest.java   | 78 +++++++++++++++++++--
 .../shenyu/protocol/mqtt/UnsubscribeTest.java      | 58 ++++++++++++++++
 13 files changed, 284 insertions(+), 92 deletions(-)

diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Connect.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Connect.java
index bca27d001d..b23b07a661 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Connect.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Connect.java
@@ -44,6 +44,12 @@ public class Connect extends MessageType {
     @Override
     public void connect(final ChannelHandlerContext ctx, final 
MqttConnectMessage msg) {
 
+        if (isConnected(ctx.channel())) {
+            LOG.info("MQTT client has already sent a CONNECT packet, closing 
connection.");
+            ctx.close().addListener(CLOSE_ON_FAILURE);
+            return;
+        }
+
         String clientId = msg.payload().clientIdentifier();
         if (StringUtils.isEmpty(clientId)) {
             LOG.info("MQTT clientId can not be empty.");
@@ -73,7 +79,7 @@ public class Connect extends MessageType {
                 .sessionPresent(true)
                 .build();
         ctx.writeAndFlush(ackMessage);
-        setConnected(true);
+        setConnected(ctx.channel(), true);
     }
 
     private void close(final ChannelHandlerContext ctx, final 
MqttConnectReturnCode returnCode) {
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MessageType.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MessageType.java
index 9d766f31bc..f6bdecac39 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MessageType.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MessageType.java
@@ -17,33 +17,37 @@
 
 package org.apache.shenyu.protocol.mqtt;
 
+import io.netty.channel.Channel;
 import io.netty.channel.ChannelHandlerContext;
 import io.netty.handler.codec.mqtt.MqttConnectMessage;
 import io.netty.handler.codec.mqtt.MqttPublishMessage;
 import io.netty.handler.codec.mqtt.MqttSubscribeMessage;
 import io.netty.handler.codec.mqtt.MqttUnsubscribeMessage;
+import io.netty.util.AttributeKey;
 
 /**
  * Command messages.
  */
 public class MessageType implements AbstractMessageType {
 
-    private volatile boolean connected;
+    private static final AttributeKey<Boolean> CONNECTED = 
AttributeKey.valueOf("connected");
 
     /**
      * isConnected.
+     * @param channel channel
      * @return connected
      */
-    boolean isConnected() {
-        return connected;
+    protected boolean isConnected(final Channel channel) {
+        return Boolean.TRUE.equals(channel.attr(CONNECTED).get());
     }
 
     /**
      * set connected.
+     * @param channel channel
      * @param connected connected
      */
-    void setConnected(final boolean connected) {
-        this.connected = connected;
+    protected void setConnected(final Channel channel, final boolean 
connected) {
+        channel.attr(CONNECTED).set(connected);
     }
 
     @Override
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java
index 49ddc8a058..a6fa975297 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandler.java
@@ -22,6 +22,8 @@ import io.netty.channel.ChannelInboundHandlerAdapter;
 import io.netty.handler.codec.mqtt.MqttMessage;
 import io.netty.util.concurrent.Future;
 import io.netty.util.concurrent.GenericFutureListener;
+import org.apache.shenyu.common.utils.Singleton;
+import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
 
 /**
  * mqtt transport handler.
@@ -38,6 +40,12 @@ public class MqttTransportHandler extends 
ChannelInboundHandlerAdapter implement
         }
     }
 
+    @Override
+    public void channelInactive(final ChannelHandlerContext ctx) throws 
Exception {
+        Singleton.INST.get(ChannelRepository.class).remove(ctx.channel());
+        ctx.fireChannelInactive();
+    }
+
     @Override
     public void operationComplete(final Future<? super Void> future) throws 
Exception {
 
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PingReq.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PingReq.java
index 716cf54915..a95db72cb4 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PingReq.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/PingReq.java
@@ -19,6 +19,8 @@ package org.apache.shenyu.protocol.mqtt;
 
 import io.netty.channel.ChannelHandlerContext;
 
+import static io.netty.channel.ChannelFutureListener.FIRE_EXCEPTION_ON_FAILURE;
+
 /**
  * Client sends pingreq to the server.
  */
@@ -26,6 +28,10 @@ public class PingReq extends MessageType {
 
     @Override
     public void pingReq(final ChannelHandlerContext ctx) {
+        if (!isConnected(ctx.channel())) {
+            ctx.channel().close().addListener(FIRE_EXCEPTION_ON_FAILURE);
+            return;
+        }
         new PingResp().pingResp(ctx);
     }
 }
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java
index bd7f48ef96..c469b3cf22 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Publish.java
@@ -35,6 +35,7 @@ import 
org.apache.shenyu.protocol.mqtt.repositories.TopicRepository;
 import java.util.List;
 import java.util.concurrent.CompletableFuture;
 
+import static io.netty.channel.ChannelFutureListener.FIRE_EXCEPTION_ON_FAILURE;
 import static io.netty.handler.codec.mqtt.MqttMessageType.PUBACK;
 
 /**
@@ -44,7 +45,8 @@ public class Publish extends MessageType {
 
     @Override
     public void publish(final ChannelHandlerContext ctx, final 
MqttPublishMessage msg) {
-        if (isConnected()) {
+        if (!isConnected(ctx.channel())) {
+            ctx.channel().close().addListener(FIRE_EXCEPTION_ON_FAILURE);
             return;
         }
         String topic = msg.variableHeader().topicName();
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Subscribe.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Subscribe.java
index c71a218ca5..1a9975cfd1 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Subscribe.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Subscribe.java
@@ -53,7 +53,7 @@ public class Subscribe extends MessageType {
     public void subscribe(final ChannelHandlerContext ctx, final 
MqttSubscribeMessage msg) {
         Channel channel = ctx.channel();
 
-        if (isConnected()) {
+        if (!isConnected(channel)) {
             channel.close().addListener(FIRE_EXCEPTION_ON_FAILURE);
             return;
         }
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Unsubscribe.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Unsubscribe.java
index 376015a6bf..154b61af72 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Unsubscribe.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/Unsubscribe.java
@@ -29,6 +29,7 @@ import 
org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository;
 
 import java.util.List;
 
+import static io.netty.channel.ChannelFutureListener.FIRE_EXCEPTION_ON_FAILURE;
 import static io.netty.handler.codec.mqtt.MqttMessageIdVariableHeader.from;
 
 /**
@@ -38,7 +39,8 @@ public class Unsubscribe extends MessageType {
 
     @Override
     public void unsubscribe(final ChannelHandlerContext ctx, final 
MqttUnsubscribeMessage msg) {
-        if (isConnected()) {
+        if (!isConnected(ctx.channel())) {
+            ctx.channel().close().addListener(FIRE_EXCEPTION_ON_FAILURE);
             return;
         }
         List<String> topics = msg.payload().topics();
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/ChannelRepository.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/ChannelRepository.java
index decd4c2e06..fe81cd1379 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/ChannelRepository.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/main/java/org/apache/shenyu/protocol/mqtt/repositories/ChannelRepository.java
@@ -20,7 +20,6 @@ package org.apache.shenyu.protocol.mqtt.repositories;
 import io.netty.channel.Channel;
 
 import java.util.Map;
-import java.util.concurrent.CompletableFuture;
 import java.util.concurrent.ConcurrentHashMap;
 
 /**
@@ -32,7 +31,7 @@ public class ChannelRepository implements 
BaseRepository<Channel, String> {
 
     @Override
     public void add(final Channel channel, final String clientId) {
-        CompletableFuture.runAsync(() -> CHANNEL_FACTORY.put(channel, 
clientId));
+        CHANNEL_FACTORY.put(channel, clientId);
     }
 
     @Override
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
index 0be443f7c9..2b8acc8e33 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
@@ -17,9 +17,9 @@
 
 package org.apache.shenyu.protocol.mqtt;
 
-import io.netty.channel.Channel;
-import io.netty.channel.ChannelFuture;
 import io.netty.channel.ChannelHandlerContext;
+import io.netty.channel.ChannelInboundHandlerAdapter;
+import io.netty.channel.embedded.EmbeddedChannel;
 import io.netty.handler.codec.mqtt.MqttConnAckMessage;
 import io.netty.handler.codec.mqtt.MqttConnectMessage;
 import io.netty.handler.codec.mqtt.MqttConnectPayload;
@@ -33,7 +33,6 @@ import 
org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
 import org.junit.jupiter.api.AfterAll;
 import org.junit.jupiter.api.BeforeAll;
 import org.junit.jupiter.api.Test;
-import org.mockito.ArgumentCaptor;
 
 import java.nio.charset.StandardCharsets;
 import java.time.Duration;
@@ -42,12 +41,10 @@ import static 
io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_ACCEP
 import static 
io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION;
 import static org.awaitility.Awaitility.await;
 import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertNotNull;
 import static org.junit.jupiter.api.Assertions.assertNull;
 import static org.junit.jupiter.api.Assertions.assertTrue;
-import static org.mockito.Mockito.mock;
-import static org.mockito.Mockito.times;
-import static org.mockito.Mockito.verify;
-import static org.mockito.Mockito.when;
 
 /**
  * Test cases for {@link Connect}.
@@ -93,32 +90,45 @@ public final class ConnectTest {
 
     @Test
     public void unsupportedProtocolVersionIsRejected() {
-        ChannelHandlerContext ctx = mock(ChannelHandlerContext.class);
-        Channel channel = mock(Channel.class);
-        when(ctx.channel()).thenReturn(channel);
-        when(ctx.close()).thenReturn(mock(ChannelFuture.class));
+        EmbeddedChannel channel = new EmbeddedChannel(new 
ChannelInboundHandlerAdapter());
+        ChannelHandlerContext ctx = channel.pipeline().lastContext();
 
         new Connect().connect(ctx, connectMessage("MQTT", 6));
 
-        ArgumentCaptor<MqttConnAckMessage> captor = 
ArgumentCaptor.forClass(MqttConnAckMessage.class);
-        verify(ctx, times(1)).writeAndFlush(captor.capture());
+        MqttConnAckMessage ackMessage = channel.readOutbound();
+        assertNotNull(ackMessage);
         assertEquals(CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION,
-                captor.getValue().variableHeader().connectReturnCode());
-        verify(ctx).close();
+                ackMessage.variableHeader().connectReturnCode());
+        channel.runPendingTasks();
+        assertFalse(channel.isActive());
         assertNull(channelRepository.get(channel));
     }
 
+    @Test
+    public void duplicateConnectIsRejected() {
+        EmbeddedChannel channel = new EmbeddedChannel(new 
ChannelInboundHandlerAdapter());
+        ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+        new Connect().connect(ctx, 
connectMessage(MqttVersion.MQTT_3_1_1.protocolName(), 
MqttVersion.MQTT_3_1_1.protocolLevel()));
+        assertNotNull(channel.readOutbound());
+
+        new Connect().connect(ctx, 
connectMessage(MqttVersion.MQTT_3_1_1.protocolName(), 
MqttVersion.MQTT_3_1_1.protocolLevel()));
+
+        channel.runPendingTasks();
+        assertFalse(channel.isActive());
+        assertNull(channel.readOutbound());
+    }
+
     private void connectIsAccepted(final MqttVersion version) {
-        ChannelHandlerContext ctx = mock(ChannelHandlerContext.class);
-        Channel channel = mock(Channel.class);
-        when(ctx.channel()).thenReturn(channel);
+        EmbeddedChannel channel = new EmbeddedChannel(new 
ChannelInboundHandlerAdapter());
+        ChannelHandlerContext ctx = channel.pipeline().lastContext();
 
         new Connect().connect(ctx, connectMessage(version.protocolName(), 
version.protocolLevel()));
 
-        ArgumentCaptor<MqttConnAckMessage> captor = 
ArgumentCaptor.forClass(MqttConnAckMessage.class);
-        verify(ctx).writeAndFlush(captor.capture());
-        assertEquals(CONNECTION_ACCEPTED, 
captor.getValue().variableHeader().connectReturnCode());
-        assertTrue(captor.getValue().variableHeader().isSessionPresent());
+        MqttConnAckMessage ackMessage = channel.readOutbound();
+        assertNotNull(ackMessage);
+        assertEquals(CONNECTION_ACCEPTED, 
ackMessage.variableHeader().connectReturnCode());
+        assertTrue(ackMessage.variableHeader().isSessionPresent());
         await().atMost(Duration.ofSeconds(5))
                 .until(() -> CLIENT_ID.equals(channelRepository.get(channel)));
     }
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java
similarity index 51%
copy from 
shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
copy to 
shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java
index 0be443f7c9..2f267af7e1 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/ConnectTest.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/MqttTransportHandlerTest.java
@@ -17,10 +17,7 @@
 
 package org.apache.shenyu.protocol.mqtt;
 
-import io.netty.channel.Channel;
-import io.netty.channel.ChannelFuture;
-import io.netty.channel.ChannelHandlerContext;
-import io.netty.handler.codec.mqtt.MqttConnAckMessage;
+import io.netty.channel.embedded.EmbeddedChannel;
 import io.netty.handler.codec.mqtt.MqttConnectMessage;
 import io.netty.handler.codec.mqtt.MqttConnectPayload;
 import io.netty.handler.codec.mqtt.MqttConnectVariableHeader;
@@ -33,26 +30,17 @@ import 
org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
 import org.junit.jupiter.api.AfterAll;
 import org.junit.jupiter.api.BeforeAll;
 import org.junit.jupiter.api.Test;
-import org.mockito.ArgumentCaptor;
 
 import java.nio.charset.StandardCharsets;
-import java.time.Duration;
 
-import static 
io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_ACCEPTED;
-import static 
io.netty.handler.codec.mqtt.MqttConnectReturnCode.CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION;
-import static org.awaitility.Awaitility.await;
 import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertFalse;
 import static org.junit.jupiter.api.Assertions.assertNull;
-import static org.junit.jupiter.api.Assertions.assertTrue;
-import static org.mockito.Mockito.mock;
-import static org.mockito.Mockito.times;
-import static org.mockito.Mockito.verify;
-import static org.mockito.Mockito.when;
 
 /**
- * Test cases for {@link Connect}.
+ * Test cases for {@link MqttTransportHandler}.
  */
-public final class ConnectTest {
+public final class MqttTransportHandlerTest {
 
     private static final String CLIENT_ID = "test-client";
 
@@ -77,58 +65,43 @@ public final class ConnectTest {
     }
 
     @Test
-    public void mqtt31ConnectIsAccepted() {
-        connectIsAccepted(MqttVersion.MQTT_3_1);
-    }
+    public void duplicateConnectCleansUpChannelRepository() {
+        EmbeddedChannel channel = new EmbeddedChannel(new 
MqttTransportHandler());
 
-    @Test
-    public void mqtt311ConnectIsAccepted() {
-        connectIsAccepted(MqttVersion.MQTT_3_1_1);
-    }
+        channel.writeInbound(connectMessage());
+        assertEquals(CLIENT_ID, channelRepository.get(channel));
 
-    @Test
-    public void mqtt5ConnectIsAccepted() {
-        connectIsAccepted(MqttVersion.MQTT_5);
-    }
+        channel.writeInbound(connectMessage());
+        channel.runPendingTasks();
 
-    @Test
-    public void unsupportedProtocolVersionIsRejected() {
-        ChannelHandlerContext ctx = mock(ChannelHandlerContext.class);
-        Channel channel = mock(Channel.class);
-        when(ctx.channel()).thenReturn(channel);
-        when(ctx.close()).thenReturn(mock(ChannelFuture.class));
-
-        new Connect().connect(ctx, connectMessage("MQTT", 6));
-
-        ArgumentCaptor<MqttConnAckMessage> captor = 
ArgumentCaptor.forClass(MqttConnAckMessage.class);
-        verify(ctx, times(1)).writeAndFlush(captor.capture());
-        assertEquals(CONNECTION_REFUSED_UNACCEPTABLE_PROTOCOL_VERSION,
-                captor.getValue().variableHeader().connectReturnCode());
-        verify(ctx).close();
+        assertFalse(channel.isActive());
         assertNull(channelRepository.get(channel));
+        channel.finishAndReleaseAll();
     }
 
-    private void connectIsAccepted(final MqttVersion version) {
-        ChannelHandlerContext ctx = mock(ChannelHandlerContext.class);
-        Channel channel = mock(Channel.class);
-        when(ctx.channel()).thenReturn(channel);
+    @Test
+    public void abruptChannelCloseCleansUpChannelRepository() {
+        EmbeddedChannel channel = new EmbeddedChannel(new 
MqttTransportHandler());
+
+        channel.writeInbound(connectMessage());
+        assertEquals(CLIENT_ID, channelRepository.get(channel));
 
-        new Connect().connect(ctx, connectMessage(version.protocolName(), 
version.protocolLevel()));
+        channel.close();
+        channel.runPendingTasks();
 
-        ArgumentCaptor<MqttConnAckMessage> captor = 
ArgumentCaptor.forClass(MqttConnAckMessage.class);
-        verify(ctx).writeAndFlush(captor.capture());
-        assertEquals(CONNECTION_ACCEPTED, 
captor.getValue().variableHeader().connectReturnCode());
-        assertTrue(captor.getValue().variableHeader().isSessionPresent());
-        await().atMost(Duration.ofSeconds(5))
-                .until(() -> CLIENT_ID.equals(channelRepository.get(channel)));
+        assertFalse(channel.isActive());
+        assertNull(channelRepository.get(channel));
+        channel.finishAndReleaseAll();
     }
 
-    private MqttConnectMessage connectMessage(final String protocolName, final 
int protocolLevel) {
+    private MqttConnectMessage connectMessage() {
         MqttFixedHeader fixedHeader = new 
MqttFixedHeader(MqttMessageType.CONNECT, false, MqttQoS.AT_MOST_ONCE, false, 0);
-        MqttConnectVariableHeader variableHeader = new 
MqttConnectVariableHeader(protocolName, protocolLevel,
+        MqttConnectVariableHeader variableHeader = new 
MqttConnectVariableHeader(
+                MqttVersion.MQTT_3_1_1.protocolName(), 
MqttVersion.MQTT_3_1_1.protocolLevel(),
                 true, true, false, 0, false, false, 60);
         MqttConnectPayload payload = new MqttConnectPayload(CLIENT_ID, null, 
null,
                 USER_NAME, PASSWORD.getBytes(StandardCharsets.UTF_8));
         return new MqttConnectMessage(fixedHeader, variableHeader, payload);
     }
+
 }
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PingReqTest.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PingReqTest.java
new file mode 100644
index 0000000000..9b26ecad9f
--- /dev/null
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PingReqTest.java
@@ -0,0 +1,56 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.shenyu.protocol.mqtt;
+
+import io.netty.channel.ChannelHandlerContext;
+import io.netty.channel.ChannelInboundHandlerAdapter;
+import io.netty.channel.embedded.EmbeddedChannel;
+import org.junit.jupiter.api.Test;
+
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertNotNull;
+import static org.junit.jupiter.api.Assertions.assertNull;
+
+/**
+ * Test cases for {@link PingReq}.
+ */
+public final class PingReqTest {
+
+    @Test
+    public void pingReqBeforeConnectClosesChannel() {
+        EmbeddedChannel channel = new EmbeddedChannel(new 
ChannelInboundHandlerAdapter());
+        ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+        new PingReq().pingReq(ctx);
+
+        channel.runPendingTasks();
+        assertFalse(channel.isActive());
+        assertNull(channel.readOutbound());
+    }
+
+    @Test
+    public void pingReqAfterConnectSendsPingResp() {
+        EmbeddedChannel channel = new EmbeddedChannel(new 
ChannelInboundHandlerAdapter());
+        new MessageType().setConnected(channel, true);
+        ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+        new PingReq().pingReq(ctx);
+
+        assertNotNull(channel.readOutbound());
+    }
+}
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishTest.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishTest.java
index 6d25ff0352..9b5d56ef01 100644
--- 
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishTest.java
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/PublishTest.java
@@ -19,23 +19,32 @@ package org.apache.shenyu.protocol.mqtt;
 
 import io.netty.buffer.Unpooled;
 import io.netty.channel.ChannelHandlerContext;
+import io.netty.channel.ChannelInboundHandlerAdapter;
+import io.netty.channel.embedded.EmbeddedChannel;
+import io.netty.handler.codec.mqtt.MqttConnectMessage;
+import io.netty.handler.codec.mqtt.MqttConnectPayload;
+import io.netty.handler.codec.mqtt.MqttConnectVariableHeader;
 import io.netty.handler.codec.mqtt.MqttFixedHeader;
 import io.netty.handler.codec.mqtt.MqttMessageType;
 import io.netty.handler.codec.mqtt.MqttPublishMessage;
 import io.netty.handler.codec.mqtt.MqttPublishVariableHeader;
 import io.netty.handler.codec.mqtt.MqttQoS;
+import io.netty.handler.codec.mqtt.MqttVersion;
 import io.netty.util.CharsetUtil;
 import org.apache.shenyu.common.utils.Singleton;
+import org.apache.shenyu.protocol.mqtt.repositories.ChannelRepository;
 import org.apache.shenyu.protocol.mqtt.repositories.SubscribeRepository;
 import org.apache.shenyu.protocol.mqtt.repositories.TopicRepository;
+import org.junit.jupiter.api.AfterAll;
 import org.junit.jupiter.api.BeforeAll;
 import org.junit.jupiter.api.Test;
 
+import java.nio.charset.StandardCharsets;
 import java.time.Duration;
 
 import static org.awaitility.Awaitility.await;
+import static org.junit.jupiter.api.Assertions.assertFalse;
 import static org.junit.jupiter.api.Assertions.assertNull;
-import static org.mockito.Mockito.mock;
 
 /**
  * Test cases for {@link Publish}.
@@ -48,6 +57,16 @@ public final class PublishTest {
 
     private static final String CLEARED_TOPIC = "test/cleared";
 
+    private static final String UNCONNECTED_TOPIC = "test/unconnected";
+
+    private static final String END_TO_END_TOPIC = "test/end-to-end";
+
+    private static final String CLIENT_ID = "test-client";
+
+    private static final String USER_NAME = "test-user";
+
+    private static final String PASSWORD = "test-password";
+
     private static TopicRepository topicRepository;
 
     @BeforeAll
@@ -55,31 +74,80 @@ public final class PublishTest {
         topicRepository = new TopicRepository();
         Singleton.INST.single(TopicRepository.class, topicRepository);
         Singleton.INST.single(SubscribeRepository.class, new 
SubscribeRepository());
+        Singleton.INST.single(ChannelRepository.class, new 
ChannelRepository());
+        new MqttContext().setUserName(USER_NAME);
+        new MqttContext().setPassword(PASSWORD);
+    }
+
+    @AfterAll
+    static void tearDown() {
+        new MqttContext().setUserName(null);
+        new MqttContext().setPassword(null);
     }
 
     @Test
     public void retainedPublishStoresMessage() {
-        new Publish().publish(mock(ChannelHandlerContext.class), 
publishMessage(RETAINED_TOPIC, "hello", true));
+        new Publish().publish(connectedContext(), 
publishMessage(RETAINED_TOPIC, "hello", true));
         await().atMost(Duration.ofSeconds(5))
                 .until(() -> 
"hello".equals(topicRepository.get(RETAINED_TOPIC)));
     }
 
     @Test
     public void nonRetainedPublishDoesNotStoreMessage() {
-        new Publish().publish(mock(ChannelHandlerContext.class), 
publishMessage(NON_RETAINED_TOPIC, "hello", false));
+        new Publish().publish(connectedContext(), 
publishMessage(NON_RETAINED_TOPIC, "hello", false));
         assertNull(topicRepository.get(NON_RETAINED_TOPIC));
     }
 
+    @Test
+    public void publishBeforeConnectClosesChannel() {
+        EmbeddedChannel channel = new EmbeddedChannel(new 
ChannelInboundHandlerAdapter());
+        ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+        new Publish().publish(ctx, publishMessage(UNCONNECTED_TOPIC, "hello", 
true));
+
+        channel.runPendingTasks();
+        assertFalse(channel.isActive());
+        assertNull(topicRepository.get(UNCONNECTED_TOPIC));
+    }
+
+    @Test
+    public void publishAfterConnectOnSameChannelIsAccepted() {
+        EmbeddedChannel channel = new EmbeddedChannel(new 
ChannelInboundHandlerAdapter());
+        ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+        new Connect().connect(ctx, connectMessage());
+        new Publish().publish(ctx, publishMessage(END_TO_END_TOPIC, "hello", 
true));
+
+        await().atMost(Duration.ofSeconds(5))
+                .until(() -> 
"hello".equals(topicRepository.get(END_TO_END_TOPIC)));
+    }
+
     @Test
     public void zeroByteRetainedPublishClearsRetainedMessage() {
         Publish publish = new Publish();
-        publish.publish(mock(ChannelHandlerContext.class), 
publishMessage(CLEARED_TOPIC, "hello", true));
+        publish.publish(connectedContext(), publishMessage(CLEARED_TOPIC, 
"hello", true));
         await().atMost(Duration.ofSeconds(5))
                 .until(() -> 
"hello".equals(topicRepository.get(CLEARED_TOPIC)));
-        publish.publish(mock(ChannelHandlerContext.class), 
publishMessage(CLEARED_TOPIC, "", true));
+        publish.publish(connectedContext(), publishMessage(CLEARED_TOPIC, "", 
true));
         assertNull(topicRepository.get(CLEARED_TOPIC));
     }
 
+    private ChannelHandlerContext connectedContext() {
+        EmbeddedChannel channel = new EmbeddedChannel(new 
ChannelInboundHandlerAdapter());
+        new MessageType().setConnected(channel, true);
+        return channel.pipeline().lastContext();
+    }
+
+    private MqttConnectMessage connectMessage() {
+        MqttFixedHeader fixedHeader = new 
MqttFixedHeader(MqttMessageType.CONNECT, false, MqttQoS.AT_MOST_ONCE, false, 0);
+        MqttConnectVariableHeader variableHeader = new 
MqttConnectVariableHeader(
+                MqttVersion.MQTT_3_1_1.protocolName(), 
MqttVersion.MQTT_3_1_1.protocolLevel(),
+                true, true, false, 0, false, false, 60);
+        MqttConnectPayload payload = new MqttConnectPayload(CLIENT_ID, null, 
null,
+                USER_NAME, PASSWORD.getBytes(StandardCharsets.UTF_8));
+        return new MqttConnectMessage(fixedHeader, variableHeader, payload);
+    }
+
     private MqttPublishMessage publishMessage(final String topic, final String 
payload, final boolean retain) {
         MqttFixedHeader fixedHeader = new 
MqttFixedHeader(MqttMessageType.PUBLISH, false, MqttQoS.AT_MOST_ONCE, retain, 
0);
         MqttPublishVariableHeader variableHeader = new 
MqttPublishVariableHeader(topic, 1);
diff --git 
a/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/UnsubscribeTest.java
 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/UnsubscribeTest.java
new file mode 100644
index 0000000000..6d942b0b65
--- /dev/null
+++ 
b/shenyu-protocol/shenyu-protocol-mqtt/src/test/java/org/apache/shenyu/protocol/mqtt/UnsubscribeTest.java
@@ -0,0 +1,58 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.shenyu.protocol.mqtt;
+
+import io.netty.channel.ChannelHandlerContext;
+import io.netty.channel.ChannelInboundHandlerAdapter;
+import io.netty.channel.embedded.EmbeddedChannel;
+import io.netty.handler.codec.mqtt.MqttFixedHeader;
+import io.netty.handler.codec.mqtt.MqttMessageIdVariableHeader;
+import io.netty.handler.codec.mqtt.MqttMessageType;
+import io.netty.handler.codec.mqtt.MqttQoS;
+import io.netty.handler.codec.mqtt.MqttUnsubscribeMessage;
+import io.netty.handler.codec.mqtt.MqttUnsubscribePayload;
+import org.junit.jupiter.api.Test;
+
+import java.util.Collections;
+
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertNull;
+
+/**
+ * Test cases for {@link Unsubscribe}.
+ */
+public final class UnsubscribeTest {
+
+    @Test
+    public void unsubscribeBeforeConnectClosesChannel() {
+        EmbeddedChannel channel = new EmbeddedChannel(new 
ChannelInboundHandlerAdapter());
+        ChannelHandlerContext ctx = channel.pipeline().lastContext();
+
+        new Unsubscribe().unsubscribe(ctx, unsubscribeMessage());
+
+        channel.runPendingTasks();
+        assertFalse(channel.isActive());
+        assertNull(channel.readOutbound());
+    }
+
+    private MqttUnsubscribeMessage unsubscribeMessage() {
+        MqttFixedHeader fixedHeader = new 
MqttFixedHeader(MqttMessageType.UNSUBSCRIBE, false, MqttQoS.AT_MOST_ONCE, 
false, 0);
+        return new MqttUnsubscribeMessage(fixedHeader, 
MqttMessageIdVariableHeader.from(1),
+                new 
MqttUnsubscribePayload(Collections.singletonList("test/topic")));
+    }
+}

Reply via email to