From 062e0deeb307df05f9f648a68d567318571baaff Mon Sep 17 00:00:00 2001 From: Gu Jiawei Date: Sun, 30 Aug 2026 22:01:17 +0800 Subject: [PATCH] 1. fix oversized packet mem leakage. --- .../mqtt/handler/MQTTPacketFilter.java | 3 ++ .../mqtt/handler/MQTTPacketFilterTest.java | 38 +++++++++++++++++++ 2 files changed, 41 insertions(+) diff --git a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/MQTTPacketFilter.java b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/MQTTPacketFilter.java index 7dc2c5be3..e7f89d880 100644 --- a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/MQTTPacketFilter.java +++ b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/MQTTPacketFilter.java @@ -40,6 +40,7 @@ import io.netty.channel.ChannelPromise; import io.netty.handler.codec.mqtt.MqttMessage; import io.netty.handler.codec.mqtt.MqttMessageType; +import io.netty.util.ReferenceCountUtil; import io.netty.util.concurrent.Future; import io.netty.util.concurrent.GenericFutureListener; import java.util.Objects; @@ -100,6 +101,8 @@ public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise promise) getLocal(OversizePacketDropped.class) .mqttPacketType(mqttMessage.fixedHeader().messageType().value()) .clientInfo(clientInfo)); + ReferenceCountUtil.release(msg); + promise.setSuccess(); } private GenericFutureListener> logMetric(MqttMessage message, int size) { diff --git a/bifromq-mqtt/bifromq-mqtt-server/src/test/java/org/apache/bifromq/mqtt/handler/MQTTPacketFilterTest.java b/bifromq-mqtt/bifromq-mqtt-server/src/test/java/org/apache/bifromq/mqtt/handler/MQTTPacketFilterTest.java index e70d448f0..e90a9357e 100644 --- a/bifromq-mqtt/bifromq-mqtt-server/src/test/java/org/apache/bifromq/mqtt/handler/MQTTPacketFilterTest.java +++ b/bifromq-mqtt/bifromq-mqtt-server/src/test/java/org/apache/bifromq/mqtt/handler/MQTTPacketFilterTest.java @@ -33,6 +33,7 @@ import static org.mockito.Mockito.when; import static org.testng.Assert.assertEquals; import static org.testng.Assert.assertNull; +import static org.testng.Assert.assertTrue; import org.apache.bifromq.metrics.ITenantMeter; import org.apache.bifromq.metrics.TenantMetric; @@ -43,14 +44,18 @@ import org.apache.bifromq.plugin.settingprovider.Setting; import org.apache.bifromq.type.ClientInfo; import io.micrometer.core.instrument.Timer; +import io.netty.buffer.ByteBuf; import io.netty.buffer.Unpooled; +import io.netty.channel.ChannelPromise; import io.netty.channel.embedded.EmbeddedChannel; import io.netty.handler.codec.mqtt.MqttMessage; import io.netty.handler.codec.mqtt.MqttMessageBuilders; import io.netty.handler.codec.mqtt.MqttProperties; import io.netty.handler.codec.mqtt.MqttPubReplyMessageVariableHeader; +import io.netty.handler.codec.mqtt.MqttPublishMessage; import io.netty.handler.codec.mqtt.MqttQoS; import java.util.List; +import java.util.concurrent.atomic.AtomicBoolean; import org.mockito.Mock; import org.mockito.MockedStatic; import org.testng.annotations.Test; @@ -201,4 +206,37 @@ public void dropUntrimableMessage() { assertNull(channel.readOutbound()); } } + + @Test + public void mqtt5DropUntrimablePublishReleasesPayloadAndCompletesPromise() { + try (MockedStatic mockedStatic = mockStatic(ITenantMeter.class)) { + mockedStatic.when(() -> ITenantMeter.get(tenantId)).thenReturn(tenantMeter); + when(tenantMeter.timer(any())).thenReturn(timer); + MQTTPacketFilter testFilter = + new MQTTPacketFilter(108, settings, mqtt5Client, eventCollector); + EmbeddedChannel channel = new EmbeddedChannel(testFilter); + MqttProperties props = new MqttProperties(); + props.add(new MqttProperties.UserProperties(List.of(new MqttProperties.StringPair("key", "val")))); + props.add(new MqttProperties.StringProperty(MqttProperties.MqttPropertyType.REASON_STRING.value(), + "11111111111")); + + MqttPublishMessage largeMessage = MqttMessageBuilders.publish() + .topicName("topic") + .qos(MqttQoS.AT_MOST_ONCE) + .payload(Unpooled.wrappedBuffer(new byte[100])) + .properties(props) + .build(); + ByteBuf payload = largeMessage.payload(); + ChannelPromise promise = channel.newPromise(); + AtomicBoolean completionListenerCalled = new AtomicBoolean(); + promise.addListener(ignored -> completionListenerCalled.set(true)); + channel.writeAndFlush(largeMessage, promise); + verify(tenantMeter, never()).recordSummary(eq(TenantMetric.MqttEgressBytes), anyDouble()); + verify(eventCollector).report(argThat(e -> e instanceof OversizePacketDropped)); + assertNull(channel.readOutbound()); + assertTrue(promise.isSuccess()); + assertTrue(completionListenerCalled.get()); + assertEquals(payload.refCnt(), 0); + } + } }