diff --git a/common/transport/mqtt/src/main/java/org/thingsboard/server/transport/mqtt/MqttTransportService.java b/common/transport/mqtt/src/main/java/org/thingsboard/server/transport/mqtt/MqttTransportService.java index 5975e6ec97..5e52863791 100644 --- a/common/transport/mqtt/src/main/java/org/thingsboard/server/transport/mqtt/MqttTransportService.java +++ b/common/transport/mqtt/src/main/java/org/thingsboard/server/transport/mqtt/MqttTransportService.java @@ -96,17 +96,28 @@ public class MqttTransportService implements TbTransportService { .childOption(ChannelOption.SO_KEEPALIVE, keepAlive); sslServerChannel = b.bind(sslHost, sslPort).sync().channel(); } - } catch (Exception e) { + } catch (Throwable e) { log.error("Failed to start MQTT transport, releasing resources", e); - if (serverChannel != null) { - serverChannel.close(); + if (e instanceof InterruptedException) { + Thread.currentThread().interrupt(); } - if (sslServerChannel != null) { - sslServerChannel.close(); + try { + if (serverChannel != null) { + serverChannel.close().sync(); + } + if (sslServerChannel != null) { + sslServerChannel.close().sync(); + } + } catch (InterruptedException ie) { + Thread.currentThread().interrupt(); + } finally { + workerGroup.shutdownGracefully(); + bossGroup.shutdownGracefully(); } - workerGroup.shutdownGracefully(); - bossGroup.shutdownGracefully(); - throw e; + if (e instanceof Exception) { + throw (Exception) e; + } + throw (Error) e; } log.info("Mqtt transport started!"); } diff --git a/common/transport/mqtt/src/test/java/org/thingsboard/server/transport/mqtt/MqttTransportServiceTest.java b/common/transport/mqtt/src/test/java/org/thingsboard/server/transport/mqtt/MqttTransportServiceTest.java index e8c5b35651..ab21209a48 100644 --- a/common/transport/mqtt/src/test/java/org/thingsboard/server/transport/mqtt/MqttTransportServiceTest.java +++ b/common/transport/mqtt/src/test/java/org/thingsboard/server/transport/mqtt/MqttTransportServiceTest.java @@ -20,9 +20,6 @@ import io.netty.channel.EventLoopGroup; 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 org.springframework.test.util.ReflectionTestUtils; import java.net.BindException; @@ -33,15 +30,12 @@ import java.util.concurrent.TimeUnit; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; import static org.awaitility.Awaitility.await; +import static org.mockito.Mockito.mock; -@ExtendWith(MockitoExtension.class) public class MqttTransportServiceTest { private static final String HOST = "127.0.0.1"; - @Mock - private MqttTransportContext context; - private MqttTransportService service; private ServerSocket occupiedSocket; private int occupiedPort; @@ -61,7 +55,7 @@ public class MqttTransportServiceTest { ReflectionTestUtils.setField(service, "bossGroupThreadCount", 1); ReflectionTestUtils.setField(service, "workerGroupThreadCount", 1); ReflectionTestUtils.setField(service, "keepAlive", true); - ReflectionTestUtils.setField(service, "context", context); + ReflectionTestUtils.setField(service, "context", mock(MqttTransportContext.class)); } @AfterEach @@ -76,13 +70,9 @@ public class MqttTransportServiceTest { assertThatThrownBy(() -> service.init()) .isInstanceOf(BindException.class); - Channel serverChannel = (Channel) ReflectionTestUtils.getField(service, "serverChannel"); - Channel sslServerChannel = (Channel) ReflectionTestUtils.getField(service, "sslServerChannel"); EventLoopGroup boss = (EventLoopGroup) ReflectionTestUtils.getField(service, "bossGroup"); EventLoopGroup worker = (EventLoopGroup) ReflectionTestUtils.getField(service, "workerGroup"); - assertThat(serverChannel).isNull(); - assertThat(sslServerChannel).isNull(); assertThat(boss).isNotNull(); assertThat(worker).isNotNull(); assertThat(boss.isShuttingDown()).isTrue();