diff --git a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/ChannelAttrs.java b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/ChannelAttrs.java index eef0de3fc..f9ecfbcf7 100644 --- a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/ChannelAttrs.java +++ b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/ChannelAttrs.java @@ -58,9 +58,25 @@ public static ChannelTrafficShapingHandler trafficShaper(ChannelHandlerContext c return ctx.channel().pipeline().get(ChannelTrafficShapingHandler.class); } + public static int maxRemainingLength(int maxPacketSize) { + if (maxPacketSize <= 2) { + return 0; + } + if (maxPacketSize <= 129) { + return maxPacketSize - 2; + } + if (maxPacketSize <= 16386) { + return maxPacketSize - 3; + } + if (maxPacketSize <= 2097155) { + return maxPacketSize - 4; + } + return maxPacketSize - 5; + } + public static void setMaxPayload(int maxUserPayloadSize, ChannelHandlerContext ctx) { ctx.channel().pipeline().replace(ctx.pipeline().get(MqttDecoder.class.getName()), MqttDecoder.class.getName(), - new MqttDecoder(maxUserPayloadSize)); + new MqttDecoder(maxRemainingLength(maxUserPayloadSize))); if (maxUserPayloadSize > ctx.channel().config().getWriteBufferHighWaterMark()) { ctx.channel().config().setWriteBufferHighWaterMark(maxUserPayloadSize + 1024); ctx.channel().config().setWriteBufferLowWaterMark(maxUserPayloadSize / 2); diff --git a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/IMQTTProtocolHelper.java b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/IMQTTProtocolHelper.java index 544277fa7..801b74092 100644 --- a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/IMQTTProtocolHelper.java +++ b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/IMQTTProtocolHelper.java @@ -90,6 +90,8 @@ public interface IMQTTProtocolHelper { int clientReceiveMaximum(); + int maxPacketSize(); + ProtocolResponse onKick(ClientInfo killer); ProtocolResponse onRedirect(boolean isPermanent, String serverReference); diff --git a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/MQTTSessionHandler.java b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/MQTTSessionHandler.java index 97c859075..b7c33e7ed 100644 --- a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/MQTTSessionHandler.java +++ b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/MQTTSessionHandler.java @@ -110,6 +110,7 @@ import org.apache.bifromq.plugin.eventcollector.Event; import org.apache.bifromq.plugin.eventcollector.IEventCollector; import org.apache.bifromq.plugin.eventcollector.OutOfTenantResource; +import org.apache.bifromq.plugin.eventcollector.mqttbroker.OversizePacketDropped; import org.apache.bifromq.plugin.eventcollector.mqttbroker.PingReq; import org.apache.bifromq.plugin.eventcollector.mqttbroker.SubStalled; import org.apache.bifromq.plugin.eventcollector.mqttbroker.accessctrl.PubActionDisallow; @@ -971,6 +972,20 @@ protected final void sendQoS0SubMessage(RoutedMessage msg) { MqttPublishMessage pubMsg = helper().buildMqttPubMessage(0, msg, false); int msgSize = sizer.sizeOf(pubMsg).encodedBytes(); assert ctx.executor().inEventLoop(); + if (msgSize > helper().maxPacketSize()) { + eventCollector.report(getLocal(OversizePacketDropped.class) + .mqttPacketType(pubMsg.fixedHeader().messageType().value()) + .clientInfo(clientInfo)); + eventCollector.report(getLocal(QoS0Dropped.class) + .reason(DropReason.ResourceExhausted) + .isRetain(msg.isRetain()) + .sender(publisher) + .topic(msg.topic()) + .matchedFilter(topicFilter) + .size(msgSize) + .clientInfo(clientInfo())); + return; + } if (!msg.permissionGranted()) { eventCollector.report(getLocal(QoS0Dropped.class) .reason(DropReason.NoSubPermission) @@ -1109,6 +1124,14 @@ private void writeConfirmableSubMessage(ConfirmingMessage confirmingMsg, boolean MqttPublishMessage pubMsg = helper().buildMqttPubMessage(packetId, msg, isDup); TopicFilterOption option = msg.option(); int msgSize = sizer.sizeOf(pubMsg).encodedBytes(); + if (msgSize > helper().maxPacketSize()) { + eventCollector.report(getLocal(OversizePacketDropped.class) + .mqttPacketType(pubMsg.fixedHeader().messageType().value()) + .clientInfo(clientInfo)); + reportDropConfirmableMsgEvent(msg, DropReason.ResourceExhausted); + ctx.executor().execute(() -> confirm(packetId, false)); + return; + } if (!msg.permissionGranted()) { reportDropConfirmableMsgEvent(msg, DropReason.NoSubPermission); ctx.executor().execute(() -> confirm(packetId, false)); diff --git a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/v3/MQTT3ProtocolHelper.java b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/v3/MQTT3ProtocolHelper.java index f9f36d990..5f6c6cdba 100644 --- a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/v3/MQTT3ProtocolHelper.java +++ b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/v3/MQTT3ProtocolHelper.java @@ -294,6 +294,11 @@ public int clientReceiveMaximum() { return 65535; } + @Override + public int maxPacketSize() { + return settings.maxPacketSize; + } + @Override public ProtocolResponse onKick(ClientInfo killer) { return goAwayNow(getLocal(Kicked.class).kicker(killer).clientInfo(clientInfo)); diff --git a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/v5/MQTT5ProtocolHelper.java b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/v5/MQTT5ProtocolHelper.java index 6cddbd1be..00fdf4423 100644 --- a/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/v5/MQTT5ProtocolHelper.java +++ b/bifromq-mqtt/bifromq-mqtt-server/src/main/java/org/apache/bifromq/mqtt/handler/v5/MQTT5ProtocolHelper.java @@ -25,6 +25,7 @@ import static org.apache.bifromq.mqtt.handler.record.ProtocolResponse.response; import static org.apache.bifromq.mqtt.handler.record.ProtocolResponse.responseNothing; import static org.apache.bifromq.mqtt.handler.v5.MQTT5MessageUtils.isUTF8Payload; +import static org.apache.bifromq.mqtt.handler.v5.MQTT5MessageUtils.maximumPacketSize; import static org.apache.bifromq.mqtt.handler.v5.MQTT5MessageUtils.messageExpiryInterval; import static org.apache.bifromq.mqtt.handler.v5.MQTT5MessageUtils.receiveMaximum; import static org.apache.bifromq.mqtt.handler.v5.MQTT5MessageUtils.requestProblemInformation; @@ -110,6 +111,7 @@ public class MQTT5ProtocolHelper implements IMQTTProtocolHelper { private final TenantSettings settings; private final ClientInfo clientInfo; private final int clientReceiveMaximum; + private final int maxPacketSize; private final boolean requestProblemInfo; private final ReceiverTopicAliasManager receiverTopicAliasManager; private final SenderTopicAliasManager senderTopicAliasManager; @@ -128,6 +130,8 @@ public MQTT5ProtocolHelper(MqttConnectMessage connMsg, Duration.ofSeconds(60)); this.clientReceiveMaximum = Math.max(settings.minSendPerSec, receiveMaximum(connMsg.variableHeader().properties()).orElse(65535)); + this.maxPacketSize = Math.min(maximumPacketSize(connMsg.variableHeader().properties()).orElse(settings.maxPacketSize), + settings.maxPacketSize); this.requestProblemInfo = requestProblemInformation(connMsg.variableHeader().properties()); } @@ -453,6 +457,11 @@ public int clientReceiveMaximum() { return clientReceiveMaximum; } + @Override + public int maxPacketSize() { + return maxPacketSize; + } + @Override public ProtocolResponse onKick(ClientInfo killer) { return farewellNow(MQTT5MessageBuilders.disconnect().reasonCode(MQTT5DisconnectReasonCode.SessionTakenOver) diff --git a/bifromq-mqtt/bifromq-mqtt-server/src/test/java/org/apache/bifromq/mqtt/handler/ChannelAttrsTest.java b/bifromq-mqtt/bifromq-mqtt-server/src/test/java/org/apache/bifromq/mqtt/handler/ChannelAttrsTest.java new file mode 100644 index 000000000..17db3ecb0 --- /dev/null +++ b/bifromq-mqtt/bifromq-mqtt-server/src/test/java/org/apache/bifromq/mqtt/handler/ChannelAttrsTest.java @@ -0,0 +1,87 @@ +/* + * 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.bifromq.mqtt.handler; + +import static org.testng.Assert.assertEquals; +import static org.testng.Assert.assertTrue; + +import org.testng.annotations.Test; + +public class ChannelAttrsTest { + + private static int varintLength(int length) { + if (length <= 127) { + return 1; + } + if (length <= 16383) { + return 2; + } + if (length <= 2097151) { + return 3; + } + return 4; + } + + private static int totalPacketSize(int remainingLength) { + return 1 + varintLength(remainingLength) + remainingLength; + } + + @Test + public void testMaxRemainingLengthBoundaries() { + assertEquals(ChannelAttrs.maxRemainingLength(0), 0); + assertEquals(ChannelAttrs.maxRemainingLength(1), 0); + assertEquals(ChannelAttrs.maxRemainingLength(2), 0); + + assertEquals(ChannelAttrs.maxRemainingLength(3), 1); + assertEquals(ChannelAttrs.maxRemainingLength(129), 127); + + assertEquals(ChannelAttrs.maxRemainingLength(130), 127); + assertEquals(ChannelAttrs.maxRemainingLength(131), 128); + assertEquals(ChannelAttrs.maxRemainingLength(16384), 16381); + assertEquals(ChannelAttrs.maxRemainingLength(16386), 16383); + + assertEquals(ChannelAttrs.maxRemainingLength(16387), 16383); + assertEquals(ChannelAttrs.maxRemainingLength(16388), 16384); + assertEquals(ChannelAttrs.maxRemainingLength(2097155), 2097151); + + assertEquals(ChannelAttrs.maxRemainingLength(2097156), 2097151); + assertEquals(ChannelAttrs.maxRemainingLength(2097157), 2097152); + } + + @Test + public void testMaxRemainingLengthExactFit() { + int[] testLimits = { + 2, 3, 4, 10, 127, 128, 129, 130, 131, 1000, 16383, 16384, 16385, 16386, 16387, 16388, + 65535, 65536, 100000, 2097154, 2097155, 2097156, 2097157, 10000000 + }; + + for (int limit : testLimits) { + int maxR = ChannelAttrs.maxRemainingLength(limit); + if (limit <= 2) { + assertEquals(maxR, 0); + } else { + assertTrue(totalPacketSize(maxR) <= limit, + "Total packet size for maxR=" + maxR + " should be <= limit=" + limit); + assertTrue(totalPacketSize(maxR + 1) > limit, + "Total packet size for (maxR+1)=" + (maxR + 1) + " should exceed limit=" + limit); + } + } + } +}