diff --git a/examples/simple_repeater/MyMesh.cpp b/examples/simple_repeater/MyMesh.cpp index ca6a3e607e..435c27c8fa 100644 --- a/examples/simple_repeater/MyMesh.cpp +++ b/examples/simple_repeater/MyMesh.cpp @@ -144,11 +144,12 @@ uint8_t MyMesh::handleLoginReq(const mesh::Identity& sender, const uint8_t* secr return 13; // reply length } -uint8_t MyMesh::handleAnonRegionsReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data) { +uint8_t MyMesh::handleAnonRegionsReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data, + size_t data_len) { if (anon_limiter.allow(rtc_clock.getCurrentTime())) { // request data has: {reply-path-len}{reply-path} + if (!mesh::Packet::hasCompletePath(data, data_len)) return 0; reply_path_len = *data++; - if (!mesh::Packet::isValidPathLen(reply_path_len)) return 0; // reject - bad encoding mesh::Packet::writePath(reply_path, data, reply_path_len); // data += (uint8_t)reply_path_len * reply_path_hash_size; @@ -162,11 +163,12 @@ uint8_t MyMesh::handleAnonRegionsReq(const mesh::Identity& sender, uint32_t send return 0; } -uint8_t MyMesh::handleAnonOwnerReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data) { +uint8_t MyMesh::handleAnonOwnerReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data, + size_t data_len) { if (anon_limiter.allow(rtc_clock.getCurrentTime())) { // request data has: {reply-path-len}{reply-path} + if (!mesh::Packet::hasCompletePath(data, data_len)) return 0; reply_path_len = *data++; - if (!mesh::Packet::isValidPathLen(reply_path_len)) return 0; // reject - bad encoding mesh::Packet::writePath(reply_path, data, reply_path_len); // data += (uint8_t)reply_path_len * reply_path_hash_size; @@ -181,11 +183,12 @@ uint8_t MyMesh::handleAnonOwnerReq(const mesh::Identity& sender, uint32_t sender return 0; } -uint8_t MyMesh::handleAnonClockReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data) { +uint8_t MyMesh::handleAnonClockReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data, + size_t data_len) { if (anon_limiter.allow(rtc_clock.getCurrentTime())) { // request data has: {reply-path-len}{reply-path} + if (!mesh::Packet::hasCompletePath(data, data_len)) return 0; reply_path_len = *data++; - if (!mesh::Packet::isValidPathLen(reply_path_len)) return 0; // reject - bad encoding mesh::Packet::writePath(reply_path, data, reply_path_len); // data += (uint8_t)reply_path_len * reply_path_hash_size; @@ -572,6 +575,8 @@ void MyMesh::onAnonDataRecv(mesh::Packet *packet, const uint8_t *secret, const m uint8_t *data, size_t len) { if (packet->getPayloadType() == PAYLOAD_TYPE_ANON_REQ) { // received an initial request by a possible admin // client (unknown at this stage) + if (len < 5 || len >= MAX_PACKET_PAYLOAD) return; + uint32_t timestamp; memcpy(×tamp, data, 4); @@ -582,11 +587,11 @@ void MyMesh::onAnonDataRecv(mesh::Packet *packet, const uint8_t *secret, const m if (data[4] == 0 || data[4] >= ' ') { // is password, ie. a login request reply_len = handleLoginReq(sender, secret, timestamp, &data[4], packet->isRouteFlood()); } else if (data[4] == ANON_REQ_TYPE_REGIONS && packet->isRouteDirect()) { - reply_len = handleAnonRegionsReq(sender, timestamp, &data[5]); + reply_len = handleAnonRegionsReq(sender, timestamp, &data[5], len - 5); } else if (data[4] == ANON_REQ_TYPE_OWNER && packet->isRouteDirect()) { - reply_len = handleAnonOwnerReq(sender, timestamp, &data[5]); + reply_len = handleAnonOwnerReq(sender, timestamp, &data[5], len - 5); } else if (data[4] == ANON_REQ_TYPE_BASIC && packet->isRouteDirect()) { - reply_len = handleAnonClockReq(sender, timestamp, &data[5]); + reply_len = handleAnonClockReq(sender, timestamp, &data[5], len - 5); } else { reply_len = 0; // unknown/invalid request type } diff --git a/examples/simple_repeater/MyMesh.h b/examples/simple_repeater/MyMesh.h index cac6c4a281..2979074c01 100644 --- a/examples/simple_repeater/MyMesh.h +++ b/examples/simple_repeater/MyMesh.h @@ -118,9 +118,9 @@ class MyMesh : public mesh::Mesh, public CommonCLICallbacks { void putNeighbour(const mesh::Identity& id, uint32_t timestamp, float snr); uint8_t handleLoginReq(const mesh::Identity& sender, const uint8_t* secret, uint32_t sender_timestamp, const uint8_t* data, bool is_flood); - uint8_t handleAnonRegionsReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data); - uint8_t handleAnonOwnerReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data); - uint8_t handleAnonClockReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data); + uint8_t handleAnonRegionsReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data, size_t data_len); + uint8_t handleAnonOwnerReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data, size_t data_len); + uint8_t handleAnonClockReq(const mesh::Identity& sender, uint32_t sender_timestamp, const uint8_t* data, size_t data_len); int handleRequest(ClientInfo* sender, uint32_t sender_timestamp, uint8_t* payload, size_t payload_len); mesh::Packet* createSelfAdvert(); diff --git a/platformio.ini b/platformio.ini index f296f83821..ddf2e8216d 100644 --- a/platformio.ini +++ b/platformio.ini @@ -172,11 +172,21 @@ build_src_filter = -<*> +<../src/Utils.cpp> +<../src/Packet.cpp> + +<../src/helpers/AdvertDataHelpers.cpp> +<../src/helpers/ConfigSerializer.cpp> +<../src/helpers/DynamicConfigSerializer.cpp> lib_deps = google/googletest @ 1.17.0 +[env:native_asan] +extends = env:native +build_unflags = -Os +build_flags = ${env:native.build_flags} + -O1 + -g + -fsanitize=address,undefined + -fno-omit-frame-pointer + [env:native_kiss_modem] platform = native test_framework = googletest diff --git a/src/Dispatcher.cpp b/src/Dispatcher.cpp index c0610b7f8a..e89660b4a5 100644 --- a/src/Dispatcher.cpp +++ b/src/Dispatcher.cpp @@ -147,45 +147,11 @@ void Dispatcher::loop() { } bool Dispatcher::tryParsePacket(Packet* pkt, const uint8_t* raw, int len) { - int i = 0; - - pkt->header = raw[i++]; - if (pkt->getPayloadVer() > PAYLOAD_VER_1) { - MESH_DEBUG_PRINTLN("%s Dispatcher::checkRecv(): unsupported packet version", getLogDateTime()); + if (len < 0 || !pkt->readFrom(raw, static_cast(len))) { + MESH_DEBUG_PRINTLN("%s Dispatcher::checkRecv(): partial, corrupt, or unsupported packet received, len=%d", getLogDateTime(), len); return false; } - - if (pkt->hasTransportCodes()) { - memcpy(&pkt->transport_codes[0], &raw[i], 2); i += 2; - memcpy(&pkt->transport_codes[1], &raw[i], 2); i += 2; - } else { - pkt->transport_codes[0] = pkt->transport_codes[1] = 0; - } - - pkt->path_len = raw[i++]; - uint8_t path_mode = pkt->path_len >> 6; // upper 2 bits (legacy firmware: 00) - if (path_mode == 3) { // Reserved for future - MESH_DEBUG_PRINTLN("%s Dispatcher::checkRecv(): unsupported path mode: 3", getLogDateTime()); - return false; - } - - uint8_t path_byte_len = (pkt->path_len & 63) * pkt->getPathHashSize(); - if (path_byte_len > MAX_PATH_SIZE || i + path_byte_len > len) { - MESH_DEBUG_PRINTLN("%s Dispatcher::checkRecv(): partial or corrupt packet received, len=%d", getLogDateTime(), len); - return false; - } - - memcpy(pkt->path, &raw[i], path_byte_len); i += path_byte_len; - - pkt->payload_len = len - i; // payload is remainder - if (pkt->payload_len > sizeof(pkt->payload)) { - MESH_DEBUG_PRINTLN("%s Dispatcher::checkRecv(): packet payload too big, payload_len=%d", getLogDateTime(), (uint32_t)pkt->payload_len); - return false; - } - - memcpy(pkt->payload, &raw[i], pkt->payload_len); - - return true; // success + return true; } void Dispatcher::checkRecv() { @@ -387,4 +353,4 @@ unsigned long Dispatcher::futureMillis(int millis_from_now) const { return _ms->getMillis() + millis_from_now; } -} \ No newline at end of file +} diff --git a/src/Mesh.cpp b/src/Mesh.cpp index c11f37cacf..ffefb8f3c9 100644 --- a/src/Mesh.cpp +++ b/src/Mesh.cpp @@ -158,12 +158,12 @@ DispatcherAction Mesh::onRecvPacket(Packet* pkt) { int len = Utils::MACThenDecrypt(secret, data, macAndData, pkt->payload_len - i); if (len > 0) { // success! if (pkt->getPayloadType() == PAYLOAD_TYPE_PATH) { + if (!Packet::isValidPathPlaintext(data, len)) { + MESH_DEBUG_PRINTLN("%s PAYLOAD_TYPE_PATH, truncated path or missing extra type", getLogDateTime()); + break; + } int k = 0; uint8_t path_len = data[k++]; - if (!Packet::isValidPathLen(path_len)) { - MESH_DEBUG_PRINTLN("%s PAYLOAD_TYPE_PATH, bad path_len: %u", getLogDateTime(), (uint32_t)path_len); - break; // reject bad encoding - } uint8_t hash_size = (path_len >> 6) + 1; uint8_t hash_count = path_len & 63; uint8_t* path = &data[k]; k += hash_size*hash_count; @@ -558,6 +558,11 @@ Packet* Mesh::createGroupDatagram(uint8_t type, const GroupChannel& channel, con } Packet* Mesh::createAck(const uint8_t* ack, uint8_t len) { + if (ack == NULL || len < MIN_ACK_PAYLOAD_SIZE || len > MAX_ACK_PAYLOAD_SIZE || + len > sizeof(Packet::payload)) { + return NULL; + } + Packet* packet = obtainNewPacket(); if (packet == NULL) { MESH_DEBUG_PRINTLN("%s Mesh::createAck(): error, packet pool empty", getLogDateTime()); @@ -572,6 +577,11 @@ Packet* Mesh::createAck(const uint8_t* ack, uint8_t len) { } Packet* Mesh::createMultiAck(const uint8_t* ack, uint8_t len, uint8_t remaining) { + if (ack == NULL || len < MIN_ACK_PAYLOAD_SIZE || len > MAX_ACK_PAYLOAD_SIZE || + len > sizeof(Packet::payload) - 1) { + return NULL; + } + Packet* packet = obtainNewPacket(); if (packet == NULL) { MESH_DEBUG_PRINTLN("%s Mesh::createMultiAck(): error, packet pool empty", getLogDateTime()); @@ -738,4 +748,4 @@ void Mesh::sendZeroHop(Packet* packet, uint16_t* transport_codes, uint32_t delay sendPacket(packet, 0, delay_millis); } -} \ No newline at end of file +} diff --git a/src/MeshCore.h b/src/MeshCore.h index 2a3582de00..990594ce5e 100644 --- a/src/MeshCore.h +++ b/src/MeshCore.h @@ -10,6 +10,13 @@ #define SEED_SIZE 32 #define SIGNATURE_SIZE 64 #define MAX_ADVERT_DATA_SIZE 32 + +// ACK payloads always start with the 4-byte acknowledgment hash. Room-server +// keep-alive ACKs append a 1-byte unsynced-message count (5 bytes total), while +// extended chat ACKs append an attempt byte and a random byte (6 bytes total). +#define MIN_ACK_PAYLOAD_SIZE 4 +#define MAX_ACK_PAYLOAD_SIZE 6 + #define CIPHER_KEY_SIZE 16 #define CIPHER_BLOCK_SIZE 16 diff --git a/src/Packet.cpp b/src/Packet.cpp index aad3e2f48e..ce4e7fac57 100644 --- a/src/Packet.cpp +++ b/src/Packet.cpp @@ -4,6 +4,60 @@ namespace mesh { +namespace { + +bool hasCompleteCiphertext(size_t payload_len, size_t prefix_len) { + const size_t overhead = prefix_len + CIPHER_MAC_SIZE; + return payload_len >= overhead + CIPHER_BLOCK_SIZE && + (payload_len - overhead) % CIPHER_BLOCK_SIZE == 0; +} + +bool isValidPayload(uint8_t header, uint8_t path_len, const uint8_t* payload, size_t payload_len) { + const uint8_t type = (header >> PH_TYPE_SHIFT) & PH_TYPE_MASK; + switch (type) { + case PAYLOAD_TYPE_REQ: + case PAYLOAD_TYPE_RESPONSE: + case PAYLOAD_TYPE_TXT_MSG: + case PAYLOAD_TYPE_PATH: + return hasCompleteCiphertext(payload_len, 2); + case PAYLOAD_TYPE_ACK: + return payload_len >= MIN_ACK_PAYLOAD_SIZE && payload_len <= MAX_ACK_PAYLOAD_SIZE; + case PAYLOAD_TYPE_ADVERT: + return payload_len >= PUB_KEY_SIZE + sizeof(uint32_t) + SIGNATURE_SIZE && + payload_len <= PUB_KEY_SIZE + sizeof(uint32_t) + SIGNATURE_SIZE + MAX_ADVERT_DATA_SIZE; + case PAYLOAD_TYPE_GRP_TXT: + case PAYLOAD_TYPE_GRP_DATA: + return hasCompleteCiphertext(payload_len, 1); + case PAYLOAD_TYPE_ANON_REQ: + return hasCompleteCiphertext(payload_len, 1 + PUB_KEY_SIZE); + case PAYLOAD_TYPE_TRACE: { + if ((path_len & 0xc0) != 0 || payload_len < 9) return false; + const uint8_t flags = payload[8]; + const size_t hash_size = 1u << (flags & 0x03); + return (flags & 0xfc) == 0 && (payload_len - 9) % hash_size == 0; + } + case PAYLOAD_TYPE_MULTIPART: + if (payload_len == 0) return false; + if ((payload[0] & 0x0f) == PAYLOAD_TYPE_ACK) { + const size_t ack_len = payload_len - 1; + return ack_len >= MIN_ACK_PAYLOAD_SIZE && ack_len <= MAX_ACK_PAYLOAD_SIZE; + } + return true; + case PAYLOAD_TYPE_CONTROL: + if (payload_len == 0) return false; + if (type == PAYLOAD_TYPE_CONTROL && (payload[0] & 0x80) != 0) { + const uint8_t route = header & PH_ROUTE_MASK; + return (route == ROUTE_TYPE_DIRECT || route == ROUTE_TYPE_TRANSPORT_DIRECT) && + (path_len & 0x3f) == 0; + } + return true; + default: + return true; + } +} + +} // namespace + Packet::Packet() { header = 0; path_len = 0; @@ -17,6 +71,18 @@ bool Packet::isValidPathLen(uint8_t path_len) { return hash_count*hash_size <= MAX_PATH_SIZE; } +bool Packet::hasCompletePath(const uint8_t* data, size_t len) { + if (data == nullptr || len == 0 || !isValidPathLen(data[0])) return false; + const size_t path_byte_len = (data[0] & 63) * ((data[0] >> 6) + 1); + return path_byte_len <= len - 1; +} + +bool Packet::isValidPathPlaintext(const uint8_t* data, size_t len) { + if (len < 2 || !hasCompletePath(data, len)) return false; + const size_t path_byte_len = (data[0] & 63) * ((data[0] >> 6) + 1); + return path_byte_len < len - 1; // One byte after the path is required for extra_type. +} + size_t Packet::writePath(uint8_t* dest, const uint8_t* src, uint8_t path_len) { uint8_t hash_count = path_len & 63; uint8_t hash_size = (path_len >> 6) + 1; @@ -62,26 +128,49 @@ uint8_t Packet::writeTo(uint8_t dest[]) const { return i; } -bool Packet::readFrom(const uint8_t src[], uint8_t len) { - uint8_t i = 0; - header = src[i++]; - if (hasTransportCodes()) { - memcpy(&transport_codes[0], &src[i], 2); i += 2; - memcpy(&transport_codes[1], &src[i], 2); i += 2; - } else { - transport_codes[0] = transport_codes[1] = 0; +bool Packet::readFrom(const uint8_t src[], size_t len) { + if (src == nullptr || len == 0 || len > MAX_TRANS_UNIT) return false; + + size_t i = 0; + const uint8_t decoded_header = src[i++]; + if (((decoded_header >> PH_VER_SHIFT) & PH_VER_MASK) > PAYLOAD_VER_1) return false; + + uint16_t decoded_transport_codes[2] = {}; + const bool has_transport_codes = + (decoded_header & PH_ROUTE_MASK) == ROUTE_TYPE_TRANSPORT_FLOOD || + (decoded_header & PH_ROUTE_MASK) == ROUTE_TYPE_TRANSPORT_DIRECT; + if (has_transport_codes) { + if (len - i < sizeof(decoded_transport_codes)) return false; + memcpy(decoded_transport_codes, &src[i], sizeof(decoded_transport_codes)); + i += sizeof(decoded_transport_codes); } - path_len = src[i++]; - if (!isValidPathLen(path_len)) return false; // bad encoding - uint8_t bl = getPathByteLen(); - memcpy(path, &src[i], bl); i += bl; + if (i >= len) return false; + const uint8_t decoded_path_len = src[i++]; + if (!isValidPathLen(decoded_path_len)) return false; - if (i >= len) return false; // bad encoding - payload_len = len - i; - if (payload_len > sizeof(payload)) return false; // bad encoding - memcpy(payload, &src[i], payload_len); //i += payload_len; + const uint8_t hash_count = decoded_path_len & 63; + const uint8_t hash_size = (decoded_path_len >> 6) + 1; + const size_t path_byte_len = hash_count * hash_size; + if (len - i < path_byte_len) return false; + + const size_t decoded_payload_len = len - i - path_byte_len; + if (decoded_payload_len > sizeof(payload)) return false; + + const uint8_t* decoded_payload = &src[i + path_byte_len]; + if (!isValidPayload(decoded_header, decoded_path_len, decoded_payload, decoded_payload_len)) { + return false; + } + + header = decoded_header; + transport_codes[0] = decoded_transport_codes[0]; + transport_codes[1] = decoded_transport_codes[1]; + path_len = decoded_path_len; + memcpy(path, &src[i], path_byte_len); + i += path_byte_len; + payload_len = decoded_payload_len; + memcpy(payload, decoded_payload, payload_len); return true; // success } -} \ No newline at end of file +} diff --git a/src/Packet.h b/src/Packet.h index c19d9e9d8f..4c50406a80 100644 --- a/src/Packet.h +++ b/src/Packet.h @@ -85,6 +85,8 @@ class Packet { static uint8_t copyPath(uint8_t* dest, const uint8_t* src, uint8_t path_len); // returns path_len static size_t writePath(uint8_t* dest, const uint8_t* src, uint8_t path_len); // returns byte length written static bool isValidPathLen(uint8_t path_len); + static bool hasCompletePath(const uint8_t* data, size_t len); + static bool isValidPathPlaintext(const uint8_t* data, size_t len); void markDoNotRetransmit() { header = 0xFF; } bool isMarkedDoNotRetransmit() const { return header == 0xFF; } @@ -108,7 +110,7 @@ class Packet { * \param src (IN) buffer containing blob * \param len the packet length (as returned by writeTo()) */ - bool readFrom(const uint8_t src[], uint8_t len); + bool readFrom(const uint8_t src[], size_t len); }; } diff --git a/src/Utils.cpp b/src/Utils.cpp index 7a3fb78b35..6c11780e73 100644 --- a/src/Utils.cpp +++ b/src/Utils.cpp @@ -51,6 +51,7 @@ void Utils::sha256(uint8_t *hash, size_t hash_len, const uint8_t* frag1, int fra } int Utils::decrypt(const uint8_t* shared_secret, uint8_t* dest, const uint8_t* src, int src_len) { + if (src_len <= 0 || src_len % CIPHER_BLOCK_SIZE != 0) return 0; #ifdef USE_CC310_HW_CRYPTO static SaSiAesUserContext_t ctx; SaSiAesUserKeyData_t keyData = { (uint8_t*)shared_secret, CIPHER_KEY_SIZE }; diff --git a/src/helpers/AdvertDataHelpers.cpp b/src/helpers/AdvertDataHelpers.cpp index 25e5cd3fe1..81fbcb92f1 100644 --- a/src/helpers/AdvertDataHelpers.cpp +++ b/src/helpers/AdvertDataHelpers.cpp @@ -1,5 +1,7 @@ #include #include +#include +#include uint8_t AdvertDataBuilder::encodeTo(uint8_t app_data[]) { app_data[0] = _type; @@ -39,33 +41,37 @@ bool AdvertDataParser::isValidName(const char *n) { AdvertDataParser::AdvertDataParser(const uint8_t app_data[], uint8_t app_data_len) { _name[0] = 0; _lat = _lon = 0; - _flags = app_data[0]; + _flags = 0; _valid = false; _extra1 = _extra2 = 0; - - int i = 1; + + if (app_data == nullptr || app_data_len == 0 || app_data_len > MAX_ADVERT_DATA_SIZE) return; + + _flags = app_data[0]; + size_t i = 1; if (_flags & ADV_LATLON_MASK) { + if (i + sizeof(_lat) + sizeof(_lon) > app_data_len) return; memcpy(&_lat, &app_data[i], 4); i += 4; memcpy(&_lon, &app_data[i], 4); i += 4; } if (_flags & ADV_FEAT1_MASK) { + if (i + sizeof(_extra1) > app_data_len) return; memcpy(&_extra1, &app_data[i], 2); i += 2; } if (_flags & ADV_FEAT2_MASK) { + if (i + sizeof(_extra2) > app_data_len) return; memcpy(&_extra2, &app_data[i], 2); i += 2; } - if (app_data_len >= i) { - int nlen = 0; - if (_flags & ADV_NAME_MASK) { - nlen = app_data_len - i; // remainder of app_data - } - if (nlen > 0) { - memcpy(_name, &app_data[i], nlen); - _name[nlen] = 0; // set null terminator - } - _valid = true; + size_t nlen = 0; + if (_flags & ADV_NAME_MASK) { + nlen = app_data_len - i; // remainder of app_data + } + if (nlen > 0) { + memcpy(_name, &app_data[i], nlen); + _name[nlen] = 0; // set null terminator } + _valid = true; } #include diff --git a/test/test_advert_data_parser/test_advert_data_parser.cpp b/test/test_advert_data_parser/test_advert_data_parser.cpp new file mode 100644 index 0000000000..0bc1fda576 --- /dev/null +++ b/test/test_advert_data_parser/test_advert_data_parser.cpp @@ -0,0 +1,91 @@ +#include + +#include +#include +#include + +#include "helpers/AdvertDataHelpers.h" + +namespace { + +void expectTruncatedInputsRejected(uint8_t flags, size_t encoded_len) { + std::vector encoded(encoded_len, 0); + encoded[0] = flags; + + for (size_t len = 0; len < encoded_len; ++len) { + const uint8_t* data = len == 0 ? nullptr : encoded.data(); + AdvertDataParser parser(data, static_cast(len)); + EXPECT_FALSE(parser.isValid()) << "accepted truncated length " << len; + } + + AdvertDataParser complete(encoded.data(), static_cast(encoded.size())); + EXPECT_TRUE(complete.isValid()); +} + +} // namespace + +TEST(AdvertDataParser, RejectsNullAndEmptyInput) { + AdvertDataParser null_input(nullptr, 0); + EXPECT_FALSE(null_input.isValid()); + + const uint8_t unused = 0; + AdvertDataParser empty_input(&unused, 0); + EXPECT_FALSE(empty_input.isValid()); +} + +TEST(AdvertDataParser, RejectsEveryTruncatedOptionalFieldCombination) { + expectTruncatedInputsRejected(ADV_LATLON_MASK, 1 + 8); + expectTruncatedInputsRejected(ADV_FEAT1_MASK, 1 + 2); + expectTruncatedInputsRejected(ADV_FEAT2_MASK, 1 + 2); + expectTruncatedInputsRejected(ADV_LATLON_MASK | ADV_FEAT1_MASK | ADV_FEAT2_MASK, + 1 + 8 + 2 + 2); +} + +TEST(AdvertDataParser, DecodesCompleteFieldsAfterBoundsChecks) { + const int32_t expected_lat = 51501234; + const int32_t expected_lon = -123456; + const uint16_t expected_feat1 = 0x1234; + const uint16_t expected_feat2 = 0xabcd; + const char expected_name[] = "repeater"; + + std::vector encoded(1 + 8 + 2 + 2 + sizeof(expected_name) - 1); + size_t offset = 0; + encoded[offset++] = ADV_LATLON_MASK | ADV_FEAT1_MASK | ADV_FEAT2_MASK | ADV_NAME_MASK; + memcpy(&encoded[offset], &expected_lat, sizeof(expected_lat)); + offset += sizeof(expected_lat); + memcpy(&encoded[offset], &expected_lon, sizeof(expected_lon)); + offset += sizeof(expected_lon); + memcpy(&encoded[offset], &expected_feat1, sizeof(expected_feat1)); + offset += sizeof(expected_feat1); + memcpy(&encoded[offset], &expected_feat2, sizeof(expected_feat2)); + offset += sizeof(expected_feat2); + memcpy(&encoded[offset], expected_name, sizeof(expected_name) - 1); + + AdvertDataParser parser(encoded.data(), static_cast(encoded.size())); + ASSERT_TRUE(parser.isValid()); + EXPECT_EQ(expected_lat, parser.getIntLat()); + EXPECT_EQ(expected_lon, parser.getIntLon()); + EXPECT_EQ(expected_feat1, parser.getFeat1()); + EXPECT_EQ(expected_feat2, parser.getFeat2()); + EXPECT_STREQ(expected_name, parser.getName()); +} + +TEST(AdvertDataParser, RejectsInputBeyondProtocolLimit) { + std::vector encoded(MAX_ADVERT_DATA_SIZE + 1, ADV_NAME_MASK); + AdvertDataParser parser(encoded.data(), static_cast(encoded.size())); + EXPECT_FALSE(parser.isValid()); +} + +TEST(AdvertDataParser, AcceptsMaximumLengthNameAndTerminatesIt) { + std::vector encoded(MAX_ADVERT_DATA_SIZE, 'a'); + encoded[0] = ADV_NAME_MASK; + + AdvertDataParser parser(encoded.data(), static_cast(encoded.size())); + ASSERT_TRUE(parser.isValid()); + EXPECT_EQ(MAX_ADVERT_DATA_SIZE - 1, strlen(parser.getName())); +} + +int main(int argc, char** argv) { + ::testing::InitGoogleTest(&argc, argv); + return RUN_ALL_TESTS(); +} diff --git a/test/test_packet_parser/test_packet_parser.cpp b/test/test_packet_parser/test_packet_parser.cpp new file mode 100644 index 0000000000..d56975aadf --- /dev/null +++ b/test/test_packet_parser/test_packet_parser.cpp @@ -0,0 +1,183 @@ +#include + +#include +#include + +#include "Packet.h" + +using mesh::Packet; + +namespace { + +bool parse(const std::vector& encoded, Packet* packet = nullptr) { + Packet local; + return (packet == nullptr ? local : *packet).readFrom(encoded.data(), encoded.size()); +} + +} // namespace + +TEST(PacketParser, RejectsTruncatedFramePrefixes) { + EXPECT_FALSE(parse({})); + EXPECT_FALSE(parse({0x0d})); // Header without path length. + + // Transport-scoped header requires two complete 16-bit transport codes and + // a path-length byte. + EXPECT_FALSE(parse({0x0c})); + EXPECT_FALSE(parse({0x0c, 0x00})); + EXPECT_FALSE(parse({0x0c, 0x00, 0x00})); + EXPECT_FALSE(parse({0x0c, 0x00, 0x00, 0x00})); + EXPECT_FALSE(parse({0x0c, 0x00, 0x00, 0x00, 0x00})); + + // One one-byte path hash is declared but absent. + EXPECT_FALSE(parse({0x0d, 0x01})); +} + +TEST(PacketParser, RejectsUnsupportedVersionAndPathMode) { + EXPECT_FALSE(parse({0x40, 0x00})); + EXPECT_FALSE(parse({0xff, 0x00})); + EXPECT_FALSE(parse({0x3e, 0xc0})); +} + +TEST(PacketParser, AllowsStructurallyValidEmptyRawPayload) { + Packet packet; + ASSERT_TRUE(parse({0x3e, 0x00}, &packet)); + EXPECT_EQ(0u, packet.path_len); + EXPECT_EQ(0u, packet.payload_len); +} + +TEST(PacketParser, RejectsTruncatedTypedPayloads) { + EXPECT_FALSE(parse({0x2e, 0x00})); // Empty CONTROL payload. + EXPECT_FALSE(parse({0x26, 0x00})); // Empty TRACE payload. + EXPECT_FALSE(parse({0x0e, 0x00, 1, 2, 3})); // Short ACK. + + // Direct encrypted payload: destination, source, MAC, then an incomplete + // AES block. A valid packet must carry at least one complete block. + std::vector direct = {0x02, 0x00, 1, 2, 3, 4, 5}; + EXPECT_FALSE(parse(direct)); + direct.resize(2 + 2 + CIPHER_MAC_SIZE + CIPHER_BLOCK_SIZE, 0); + EXPECT_TRUE(parse(direct)); +} + +TEST(PacketParser, EnforcesSupportedAckLengths) { + EXPECT_TRUE(parse({0x0e, 0x00, 1, 2, 3, 4})); + EXPECT_TRUE(parse({0x0e, 0x00, 1, 2, 3, 4, 5})); + EXPECT_TRUE(parse({0x0e, 0x00, 1, 2, 3, 4, 5, 6})); + EXPECT_FALSE(parse({0x0e, 0x00, 1, 2, 3, 4, 5, 6, 7})); + + // Multipart ACKs reserve payload[0] for their sequence/type byte. + EXPECT_TRUE(parse({0x2a, 0x00, PAYLOAD_TYPE_ACK, 1, 2, 3, 4})); + EXPECT_TRUE(parse({0x2a, 0x00, PAYLOAD_TYPE_ACK, 1, 2, 3, 4, 5, 6})); + EXPECT_FALSE(parse({0x2a, 0x00, PAYLOAD_TYPE_ACK, 1, 2, 3, 4, 5, 6, 7})); +} + +TEST(PacketParser, RejectsInvalidTraceShapeAndZeroHopControlRoute) { + std::vector trace(2 + 9, 0); + trace[0] = 0x26; + EXPECT_TRUE(parse(trace)); + + trace[10] = 0x04; // Reserved trace flag bit. + EXPECT_FALSE(parse(trace)); + trace[10] = 0x01; // Two-byte hashes, but one trailing hash byte. + trace.push_back(0xaa); + EXPECT_FALSE(parse(trace)); + + EXPECT_TRUE(parse({0x2e, 0x00, 0x80})); + EXPECT_FALSE(parse({0x2d, 0x00, 0x80})); + EXPECT_FALSE(parse({0x2e, 0x01, 0xaa, 0x80})); +} + +TEST(PacketParser, RejectsOverlongAdvertApplicationData) { + std::vector advert(2 + PUB_KEY_SIZE + sizeof(uint32_t) + SIGNATURE_SIZE, + 0); + advert[0] = 0x12; + EXPECT_TRUE(parse(advert)); + + advert.resize(advert.size() + MAX_ADVERT_DATA_SIZE, 0); + EXPECT_TRUE(parse(advert)); + + advert.push_back(0); + EXPECT_FALSE(parse(advert)); +} + +TEST(PacketParser, RejectsTruncatedDecryptedPathPlaintext) { + const uint8_t missing_path[] = {0x3f}; + EXPECT_FALSE(Packet::isValidPathPlaintext(missing_path, sizeof(missing_path))); + + const uint8_t missing_extra_type[] = {0x00}; + EXPECT_FALSE(Packet::isValidPathPlaintext(missing_extra_type, + sizeof(missing_extra_type))); + + const uint8_t valid[] = {0x01, 0xaa, 0x00}; + EXPECT_TRUE(Packet::isValidPathPlaintext(valid, sizeof(valid))); +} + +TEST(PacketParser, ValidatesLengthPrefixedPathFields) { + EXPECT_FALSE(Packet::hasCompletePath(nullptr, 0)); + + const uint8_t empty_path[] = {0x00}; + EXPECT_TRUE(Packet::hasCompletePath(empty_path, sizeof(empty_path))); + + const uint8_t truncated_path[] = {0x01}; + EXPECT_FALSE(Packet::hasCompletePath(truncated_path, sizeof(truncated_path))); + + const uint8_t complete_path[] = {0x01, 0xaa}; + EXPECT_TRUE(Packet::hasCompletePath(complete_path, sizeof(complete_path))); + + // Callers may pass authenticated AES padding after the complete path. + const uint8_t path_with_trailing_data[] = {0x01, 0xaa, 0x00}; + EXPECT_TRUE(Packet::hasCompletePath(path_with_trailing_data, + sizeof(path_with_trailing_data))); + + const uint8_t reserved_path_mode[] = {0xc0}; + EXPECT_FALSE(Packet::hasCompletePath(reserved_path_mode, + sizeof(reserved_path_mode))); + + std::vector maximum_path(1 + MAX_PATH_SIZE, 0xaa); + maximum_path[0] = 0x60; // 32 two-byte path hashes. + EXPECT_TRUE(Packet::hasCompletePath(maximum_path.data(), maximum_path.size())); + maximum_path.pop_back(); + EXPECT_FALSE(Packet::hasCompletePath(maximum_path.data(), maximum_path.size())); +} + +TEST(PacketParser, SafelyRejectsArbitraryFramesAtEveryWireLength) { + uint32_t state = 0x6d657368; + for (size_t len = 0; len <= MAX_TRANS_UNIT; ++len) { + for (int sample = 0; sample < 32; ++sample) { + std::vector input(len); + for (uint8_t& byte : input) { + state = state * 1664525u + 1013904223u; + byte = static_cast(state >> 24); + } + Packet packet; + (void)packet.readFrom(input.data(), input.size()); + } + } +} + +TEST(PacketParser, EnforcesPathAndPayloadBounds) { + // 63 two-byte hashes exceed MAX_PATH_SIZE. + EXPECT_FALSE(parse({0x3e, 0x7f})); + + std::vector maximum = {0x3e, 0x00}; + maximum.resize(2 + MAX_PACKET_PAYLOAD, 0xaa); + EXPECT_TRUE(parse(maximum)); + + maximum.push_back(0xaa); + EXPECT_FALSE(parse(maximum)); +} + +TEST(PacketParser, AcceptsCompleteTransportScopedPacket) { + Packet packet; + ASSERT_TRUE(parse({0x0f, 0x34, 0x12, 0x78, 0x56, 0x01, 0xaa, + 0x01, 0x02, 0x03, 0x04}, + &packet)); + EXPECT_EQ(0x1234u, packet.transport_codes[0]); + EXPECT_EQ(0x5678u, packet.transport_codes[1]); + EXPECT_EQ(1u, packet.getPathByteLen()); + EXPECT_EQ(4u, packet.payload_len); +} + +int main(int argc, char** argv) { + ::testing::InitGoogleTest(&argc, argv); + return RUN_ALL_TESTS(); +} diff --git a/test/test_utils/test_tohex.cpp b/test/test_utils/test_tohex.cpp index fec3ae4877..149eb76c61 100644 --- a/test/test_utils/test_tohex.cpp +++ b/test/test_utils/test_tohex.cpp @@ -51,6 +51,18 @@ TEST(UtilsToHex, NullTerminatesOnEmptyInput) { EXPECT_EQ('\0', output[0]); } +TEST(UtilsDecrypt, RejectsEmptyAndPartialCiphertextBlocks) { + uint8_t key[CIPHER_KEY_SIZE] = {}; + uint8_t input[CIPHER_BLOCK_SIZE] = {}; + uint8_t output[CIPHER_BLOCK_SIZE] = {}; + + EXPECT_EQ(0, Utils::decrypt(key, output, input, 0)); + EXPECT_EQ(0, Utils::decrypt(key, output, input, 1)); + EXPECT_EQ(0, Utils::decrypt(key, output, input, CIPHER_BLOCK_SIZE - 1)); + EXPECT_EQ(CIPHER_BLOCK_SIZE, + Utils::decrypt(key, output, input, CIPHER_BLOCK_SIZE)); +} + int main(int argc, char **argv) { ::testing::InitGoogleTest(&argc, argv); return RUN_ALL_TESTS();