diff --git a/src/iot/mqtt/server/Mqtt.cpp b/src/iot/mqtt/server/Mqtt.cpp index 0ff1afbb8a..7280a61aba 100644 --- a/src/iot/mqtt/server/Mqtt.cpp +++ b/src/iot/mqtt/server/Mqtt.cpp @@ -253,6 +253,11 @@ namespace iot::mqtt::server { << connectionName << " MQTT Broker: Wrong Protocol Level: " << MQTT_VERSION_3_1_1 << " != " << connect.getLevel(); sendConnack(MQTT_CONNACK_UNACEPTABLEVERSION, MQTT_SESSION_NEW); + mqttContext->close(); + } else if ((connect.getConnectFlags() & 0x01) != 0) { + iot::mqtt::semantic::mqttServerLog().error() + << connectionName << " MQTT Broker: CONNECT reserved flag bit set"; + mqttContext->close(); } else if (connect.isFakedClientId() && !connect.getCleanSession()) { iot::mqtt::semantic::mqttServerLog().error() << connectionName << " MQTT Broker: Resume session but no ClientId present"; diff --git a/src/iot/mqtt/server/packets/Publish.cpp b/src/iot/mqtt/server/packets/Publish.cpp index 2c12a73fa0..9e36738712 100644 --- a/src/iot/mqtt/server/packets/Publish.cpp +++ b/src/iot/mqtt/server/packets/Publish.cpp @@ -56,7 +56,7 @@ namespace iot::mqtt::server::packets { this->dup = (flags & 0x08) != 0; this->retain = (flags & 0x01) != 0; - error = this->qoS > 2; + error = this->qoS > 2 || (this->qoS == 0 && this->dup); } std::size_t Publish::deserializeVP(iot::mqtt::MqttContext* mqttContext) { @@ -69,6 +69,11 @@ namespace iot::mqtt::server::packets { break; } + if (topic.size() == 0) { + error = true; + break; + } + state++; [[fallthrough]]; case 1: @@ -77,9 +82,18 @@ namespace iot::mqtt::server::packets { if (!packetIdentifier.isComplete()) { break; } + if (packetIdentifier == 0) { + error = true; + break; + } + } + + if (getConsumed() + consumed > getRemainingLength()) { + error = true; + break; } - message.setSize(static_cast(getRemainingLength() - getConsumed() - consumed)); + message.setSize(getRemainingLength() - getConsumed() - consumed); state++; [[fallthrough]]; diff --git a/src/iot/mqtt/server/packets/Subscribe.cpp b/src/iot/mqtt/server/packets/Subscribe.cpp index bc5cfdc9e5..da53da866b 100644 --- a/src/iot/mqtt/server/packets/Subscribe.cpp +++ b/src/iot/mqtt/server/packets/Subscribe.cpp @@ -84,6 +84,10 @@ namespace iot::mqtt::server::packets { if (!qoS.isComplete()) { break; } + if (static_cast(qoS) > 2) { + error = true; + break; + } topics.emplace_back(topic, qoS); topic.reset(); qoS.reset(); diff --git a/src/iot/mqtt/types/UIntV.cpp b/src/iot/mqtt/types/UIntV.cpp index 2ddfecdd30..e4dfc647a3 100644 --- a/src/iot/mqtt/types/UIntV.cpp +++ b/src/iot/mqtt/types/UIntV.cpp @@ -69,10 +69,10 @@ namespace iot::mqtt::types { if (consumed > 0) { value.push_back(byte); - if (value.size() > sizeof(uint32_t)) { + complete = (byte & 0x80) == 0; + if (value.size() > sizeof(uint32_t) || (value.size() == sizeof(uint32_t) && !complete)) { error = true; - } else { - complete = (byte & 0x80) == 0; + complete = false; } } } while (consumed > 0 && !complete && !error); diff --git a/tests/unit/CMakeLists.txt b/tests/unit/CMakeLists.txt index 2be7f82148..d1f542bb54 100644 --- a/tests/unit/CMakeLists.txt +++ b/tests/unit/CMakeLists.txt @@ -46,4 +46,6 @@ add_subdirectory(http) add_subdirectory(log) +add_subdirectory(mqtt) + add_subdirectory(utils) diff --git a/tests/unit/mqtt/CMakeLists.txt b/tests/unit/mqtt/CMakeLists.txt new file mode 100644 index 0000000000..59ce87599e --- /dev/null +++ b/tests/unit/mqtt/CMakeLists.txt @@ -0,0 +1,5 @@ +snodec_add_test(Mqtt311PacketValidationTest Mqtt311PacketValidationTest.cpp) +target_include_directories(Mqtt311PacketValidationTest PRIVATE ${PROJECT_SOURCE_DIR}) +target_compile_features(Mqtt311PacketValidationTest PRIVATE cxx_std_20) +target_link_libraries(Mqtt311PacketValidationTest PRIVATE snodec-test-support snodec::mqtt snodec::mqtt-packets snodec::mqtt-server snodec::mqtt-server-packets) +set_tests_properties(Mqtt311PacketValidationTest PROPERTIES LABELS "unit;mqtt;mqtt311;validation" TIMEOUT 10) diff --git a/tests/unit/mqtt/Mqtt311PacketValidationTest.cpp b/tests/unit/mqtt/Mqtt311PacketValidationTest.cpp new file mode 100644 index 0000000000..ca82972ab4 --- /dev/null +++ b/tests/unit/mqtt/Mqtt311PacketValidationTest.cpp @@ -0,0 +1,218 @@ +#include "iot/mqtt/MqttContext.h" +#include "iot/mqtt/Mqtt.h" +#include "iot/mqtt/ControlPacketDeserializer.h" +#include "iot/mqtt/packets/Publish.h" +#include "iot/mqtt/FixedHeader.h" +#include "iot/mqtt/server/packets/Connect.h" +#include "iot/mqtt/server/packets/Publish.h" +#include "iot/mqtt/server/packets/Pubrel.h" +#include "iot/mqtt/server/packets/Subscribe.h" +#include "iot/mqtt/server/packets/Unsubscribe.h" +#include "tests/support/TestResult.h" + +#include +#include +#include +#include + +namespace { + +class DummyMqtt : public iot::mqtt::Mqtt { +public: + DummyMqtt() : iot::mqtt::Mqtt("test") {} + bool onSignal(int) override { return true; } + iot::mqtt::ControlPacketDeserializer* createControlPacketDeserializer(iot::mqtt::FixedHeader&) const override { return nullptr; } + void deliverPacket(iot::mqtt::ControlPacketDeserializer*) override {} + void distributePublish(const iot::mqtt::packets::Publish&) override {} +}; + +class BufferContext : public iot::mqtt::MqttContext { +public: + explicit BufferContext(const std::vector& input) + : iot::mqtt::MqttContext(new DummyMqtt()) + , input(input) { + } + + std::size_t recv(char* chunk, std::size_t chunklen) override { + const std::size_t available = input.size() - offset; + const std::size_t count = std::min(chunklen, available); + std::copy_n(input.data() + offset, count, chunk); + offset += count; + return count; + } + + void send(const char*, std::size_t) override {} + core::socket::stream::SocketConnection* getSocketConnection() const override { return nullptr; } + void end() override {} + void close() override { closed = true; } + + bool closed = false; + +private: + std::vector input; + std::size_t offset = 0; +}; + +void appendString(std::vector& data, const std::string& value) { + data.push_back(static_cast((value.size() >> 8) & 0xff)); + data.push_back(static_cast(value.size() & 0xff)); + data.insert(data.end(), value.begin(), value.end()); +} + +std::vector connectVp(uint8_t level, uint8_t flags) { + std::vector data; + appendString(data, "MQTT"); + data.push_back(static_cast(level)); + data.push_back(static_cast(flags)); + data.push_back(0); + data.push_back(60); + appendString(data, "client"); + return data; +} + +std::vector subscribeVp(std::initializer_list> topics) { + std::vector data{0, 1}; + for (const auto& [topic, qos] : topics) { + appendString(data, topic); + data.push_back(static_cast(qos)); + } + return data; +} + +std::vector publishVp(const std::string& topic, const std::string& payload, bool includePacketId = false, uint16_t pid = 1) { + std::vector data; + appendString(data, topic); + if (includePacketId) { + data.push_back(static_cast((pid >> 8) & 0xff)); + data.push_back(static_cast(pid & 0xff)); + } + data.insert(data.end(), payload.begin(), payload.end()); + return data; +} + +template +bool deserializePacket(Packet& packet, const std::vector& data) { + BufferContext context(data); + while (!packet.isComplete() && !packet.isError()) { + if (packet.deserialize(&context) == 0) { + break; + } + } + return packet.isComplete() && !packet.isError(); +} + +} // namespace + +int main() { + tests::support::TestResult result; + + { + iot::mqtt::packets::Connect normal("client", 60, true, "", "", 0, false, "", "", false); + const auto wire = normal.serialize(); + result.expectEqual(0x04, static_cast(wire[8]), "client CONNECT without loop prevention sends level 0x04"); + } + { + iot::mqtt::packets::Connect bridge("client", 60, true, "", "", 0, false, "", "", true); + const auto wire = bridge.serialize(); + result.expectEqual(0x84, static_cast(wire[8]), "client CONNECT with loop prevention sends private bridge level 0x84"); + } + + { + const auto data = connectVp(0x04, 0x02); + iot::mqtt::server::packets::Connect connect(data.size(), 0x00); + result.expectTrue(deserializePacket(connect, data), "normal MQTT 3.1.1 CONNECT parses"); + result.expectEqual(0x04, connect.getLevel(), "normal CONNECT level remains 0x04"); + result.expectTrue(connect.getReflect(), "normal CONNECT reflects publications"); + } + { + const auto data = connectVp(0x84, 0x02); + iot::mqtt::server::packets::Connect connect(data.size(), 0x00); + result.expectTrue(deserializePacket(connect, data), "private bridge CONNECT parses"); + result.expectEqual(0x04, connect.getLevel(), "private bridge CONNECT masks level to 0x04"); + result.expectTrue(!connect.getReflect(), "private bridge CONNECT disables reflection"); + } + { + const auto data = connectVp(0x84, 0x03); + iot::mqtt::server::packets::Connect connect(data.size(), 0x00); + result.expectTrue(deserializePacket(connect, data), "CONNECT parser preserves reserved flag for server validation"); + result.expectEqual(0x03, connect.getConnectFlags(), "CONNECT reserved flag remains visible"); + } + + for (uint8_t qos = 0; qos <= 2; ++qos) { + const auto data = subscribeVp({{"a/b", qos}}); + iot::mqtt::server::packets::Subscribe subscribe(data.size(), 0x02); + result.expectTrue(deserializePacket(subscribe, data), "SUBSCRIBE QoS 0/1/2 parses"); + } + for (uint8_t qos : {uint8_t{0x03}, uint8_t{0x08}, uint8_t{0x0c}, uint8_t{0x80}}) { + const auto data = subscribeVp({{"a/b", qos}}); + iot::mqtt::server::packets::Subscribe subscribe(data.size(), 0x02); + result.expectTrue(!deserializePacket(subscribe, data), "SUBSCRIBE invalid QoS/options rejected"); + } + { + const auto data = subscribeVp({{"a/b", 0}, {"c/d", 0x80}}); + iot::mqtt::server::packets::Subscribe subscribe(data.size(), 0x02); + result.expectTrue(!deserializePacket(subscribe, data), "SUBSCRIBE rejects invalid QoS on second topic"); + } + { + iot::mqtt::server::packets::Subscribe subscribe(5, 0x00); + result.expectTrue(subscribe.isError(), "SUBSCRIBE fixed-header flags other than 0x02 rejected"); + iot::mqtt::server::packets::Unsubscribe unsubscribe(5, 0x00); + result.expectTrue(unsubscribe.isError(), "UNSUBSCRIBE fixed-header flags other than 0x02 rejected"); + iot::mqtt::server::packets::Pubrel pubrel(2, 0x00); + result.expectTrue(pubrel.isError(), "PUBREL fixed-header flags other than 0x02 rejected"); + } + + { + const auto data = publishVp("a/b", "payload"); + iot::mqtt::server::packets::Publish publish(data.size(), 0x00); + result.expectTrue(deserializePacket(publish, data), "PUBLISH QoS0 DUP0 parses"); + result.expectTrue(!publish.getRetain(), "PUBLISH retain false preserved"); + } + { + const auto data = publishVp("a/b", "payload"); + iot::mqtt::server::packets::Publish publish(data.size(), 0x01); + result.expectTrue(deserializePacket(publish, data), "PUBLISH retain true parses"); + result.expectTrue(publish.getRetain(), "PUBLISH retain true preserved"); + } + { + const auto data = publishVp("a/b", "payload"); + iot::mqtt::server::packets::Publish publish(data.size(), 0x08); + result.expectTrue(publish.isError(), "PUBLISH QoS0 DUP1 rejected at fixed header"); + } + { + const auto data = publishVp("a/b", "payload", true, 1); + iot::mqtt::server::packets::Publish publish(data.size(), 0x06); + result.expectTrue(publish.isError(), "PUBLISH QoS3 rejected at fixed header"); + } + { + const auto data = publishVp("a/b", "payload", true, 0); + iot::mqtt::server::packets::Publish publish(data.size(), 0x02); + result.expectTrue(!deserializePacket(publish, data), "PUBLISH QoS1 packet id zero rejected"); + } + { + const auto data = publishVp("", "payload"); + iot::mqtt::server::packets::Publish publish(data.size(), 0x00); + result.expectTrue(!deserializePacket(publish, data), "PUBLISH empty topic rejected"); + } + { + const auto data = publishVp("a/b", std::string(70000, 'x')); + iot::mqtt::server::packets::Publish publish(data.size(), 0x00); + result.expectTrue(deserializePacket(publish, data), "PUBLISH payload larger than 65535 is not truncated"); + result.expectEqual(70000, static_cast(publish.getMessage().size()), "large PUBLISH payload size preserved"); + } + + { + BufferContext context({char(0x30), char(0xff), char(0xff), char(0xff), char(0xff), char(0x7f)}); + iot::mqtt::FixedHeader fixedHeader; + fixedHeader.deserialize(&context); + result.expectTrue(fixedHeader.isError(), "Remaining Length overlong sequence rejected"); + } + { + BufferContext context({char(0x30), char(0xff), char(0xff), char(0xff), char(0xff)}); + iot::mqtt::FixedHeader fixedHeader; + fixedHeader.deserialize(&context); + result.expectTrue(fixedHeader.isError(), "Remaining Length continuation on fourth byte rejected"); + } + + return result.processResult(); +}