diff --git a/src/kv/generic_serialise_wrapper.h b/src/kv/generic_serialise_wrapper.h index 13e5645fc16..b2a128f4d2a 100644 --- a/src/kv/generic_serialise_wrapper.h +++ b/src/kv/generic_serialise_wrapper.h @@ -347,7 +347,14 @@ namespace ccf::kv } serialized::skip(data_, size_, crypto_util->get_header_length()); - auto public_domain_length = serialized::read(data_, size_); + const auto public_domain_length = serialized::read(data_, size_); + if (public_domain_length > size_) + { + throw std::logic_error(fmt::format( + "Public domain length {} exceeds remaining entry size {}", + public_domain_length, + size_)); + } const auto* data_public = data_; public_reader.init(data_public, public_domain_length); diff --git a/src/kv/raw_serialise.h b/src/kv/raw_serialise.h index ef1b1104a2a..74bd41dd987 100644 --- a/src/kv/raw_serialise.h +++ b/src/kv/raw_serialise.h @@ -6,7 +6,11 @@ #include "generic_serialise_wrapper.h" #include +#include +#include +#include #include +#include #include namespace ccf::kv @@ -126,21 +130,40 @@ namespace ccf::kv class RawReader { - public: - const uint8_t* data_ptr; + private: + [[nodiscard]] size_t remaining_bytes() const + { + assert(data_offset <= data_size); + return data_size - data_offset; + } + + void require_bytes(size_t required, const char* description) const + { + const auto remaining = remaining_bytes(); + if (required > remaining) + { + throw std::runtime_error(fmt::format( + "Expected {} bytes for {}, found only {}", + required, + description, + remaining)); + } + } + + const uint8_t* data_ptr{nullptr}; size_t data_offset{0}; - size_t data_size; + size_t data_size{0}; + public: /** Reads the next entry, advancing data_offset */ template T read_entry() { - auto remainder = data_size - data_offset; - const auto* data = data_ptr + data_offset; - const auto entry = serialized::read(data, remainder); - const auto bytes_read = data_size - data_offset - remainder; - data_offset += bytes_read; + require_bytes(sizeof(T), "fixed-size entry"); + T entry; + std::memcpy(&entry, data_ptr + data_offset, sizeof(T)); + data_offset += sizeof(T); return entry; } @@ -148,14 +171,8 @@ namespace ccf::kv */ size_t read_size_prefixed_entry(size_t& start_offset) { - auto remainder = data_size - data_offset; - auto entry_size = read_entry(); - - if (remainder < entry_size) - { - throw std::runtime_error(fmt::format( - "Expected {} byte entry, found only {}", entry_size, remainder)); - } + const auto entry_size = read_entry(); + require_bytes(entry_size, "size-prefixed entry"); start_offset = data_offset; data_offset += entry_size; @@ -166,13 +183,18 @@ namespace ccf::kv RawReader(const RawReader& other) = delete; RawReader& operator=(const RawReader& other) = delete; - RawReader(const uint8_t* data_in_ptr = nullptr, size_t data_in_size = 0) : - data_ptr(data_in_ptr), - data_size(data_in_size) - {} + RawReader(const uint8_t* data_in_ptr = nullptr, size_t data_in_size = 0) + { + init(data_in_ptr, data_in_size); + } void init(const uint8_t* data_in_ptr, size_t data_in_size) { + if (data_in_ptr == nullptr && data_in_size != 0) + { + throw std::invalid_argument( + "Cannot initialise raw reader with null data and non-zero size"); + } data_offset = 0; data_ptr = data_in_ptr; data_size = data_in_size; @@ -186,14 +208,31 @@ namespace ccf::kv std::is_same_v) { size_t entry_offset = 0; - size_t entry_size = read_size_prefixed_entry(entry_offset); + const auto entry_size = read_size_prefixed_entry(entry_offset); + using Element = typename T::value_type; + if (entry_size % sizeof(Element) != 0) + { + throw std::runtime_error(fmt::format( + "Size-prefixed entry of {} bytes is not divisible by element size " + "{}", + entry_size, + sizeof(Element))); + } - T ret(entry_size / sizeof(typename T::value_type)); - auto* data_dest = reinterpret_cast(ret.data()); - auto capacity = entry_size; - // NOLINTNEXTLINE(readability-suspicious-call-argument) - serialized::write( - data_dest, capacity, data_ptr + entry_offset, entry_size); + T ret; + const auto element_count = entry_size / sizeof(Element); + if (element_count > ret.max_size()) + { + throw std::length_error(fmt::format( + "Size-prefixed entry contains too many elements ({})", + element_count)); + } + ret.resize(element_count); + if (entry_size == 0) + { + return ret; + } + std::memcpy(ret.data(), data_ptr + entry_offset, entry_size); return ret; } @@ -201,10 +240,17 @@ namespace ccf::kv { T ret{}; auto* data_ = reinterpret_cast(ret.data()); - constexpr size_t size = ret.size() * sizeof(typename T::value_type); - auto size_ = size; - serialized::write(data_, size_, data_ptr + data_offset, size); - data_offset += size; + constexpr auto element_count = std::tuple_size_v; + static_assert( + element_count <= + std::numeric_limits::max() / sizeof(typename T::value_type)); + constexpr size_t size = element_count * sizeof(typename T::value_type); + require_bytes(size, "fixed-size array"); + if constexpr (size > 0) + { + std::memcpy(data_, data_ptr + data_offset, size); + data_offset += size; + } return ret; } @@ -240,7 +286,8 @@ namespace ccf::kv [[nodiscard]] bool is_eos() const { - return data_offset >= data_size; + assert(data_offset <= data_size); + return data_offset == data_size; } }; diff --git a/src/kv/test/kv_serialisation.cpp b/src/kv/test/kv_serialisation.cpp index 92078bc2478..c41d1949482 100644 --- a/src/kv/test/kv_serialisation.cpp +++ b/src/kv/test/kv_serialisation.cpp @@ -2,12 +2,16 @@ // Licensed under the Apache 2.0 License. #include "ds/internal_logger.h" #include "kv/kv_serialiser.h" +#include "kv/raw_serialise.h" #include "kv/store.h" #include "kv/test/null_encryptor.h" #include "kv/test/stub_consensus.h" #include #undef FAIL +#include +#include +#include #include #include @@ -19,6 +23,138 @@ struct MapTypes using StringNum = ccf::kv::Map; }; +static std::vector make_size_prefixed_bytes( + size_t declared_size, size_t actual_size) +{ + std::vector bytes(sizeof(size_t) + actual_size); + std::memcpy(bytes.data(), &declared_size, sizeof(declared_size)); + return bytes; +} + +TEST_CASE( + "Raw reader rejects truncated entries" * doctest::test_suite("serialisation")) +{ + SUBCASE("Fixed-size integral") + { + std::vector bytes(sizeof(uint64_t) - 1); + ccf::kv::RawReader reader(bytes.data(), bytes.size()); + REQUIRE_THROWS(reader.read_next()); + } + + SUBCASE("Fixed-size array") + { + std::vector bytes(ccf::crypto::Sha256Hash::SIZE - 1); + ccf::kv::RawReader reader(bytes.data(), bytes.size()); + REQUIRE_THROWS(reader.read_next()); + } + + SUBCASE("Short size prefix") + { + std::vector bytes(sizeof(size_t) - 1); + ccf::kv::RawReader reader(bytes.data(), bytes.size()); + REQUIRE_THROWS(reader.read_next>()); + } + + SUBCASE("Prefix without payload") + { + auto bytes = make_size_prefixed_bytes(1, 0); + ccf::kv::RawReader reader(bytes.data(), bytes.size()); + REQUIRE_THROWS(reader.read_next>()); + } + + SUBCASE("Nine bytes remaining declare two payload bytes") + { + auto bytes = make_size_prefixed_bytes(2, 1); + ccf::kv::RawReader reader(bytes.data(), bytes.size()); + REQUIRE_THROWS(reader.read_next>()); + } + + SUBCASE("Impossible payload length") + { + auto bytes = + make_size_prefixed_bytes(std::numeric_limits::max(), 0); + ccf::kv::RawReader reader(bytes.data(), bytes.size()); + REQUIRE_THROWS(reader.read_next>()); + } + + SUBCASE("Payload is not a whole number of elements") + { + auto bytes = make_size_prefixed_bytes(1, 1); + ccf::kv::RawReader reader(bytes.data(), bytes.size()); + REQUIRE_THROWS(reader.read_next>()); + } + + SUBCASE("Null data with non-zero size") + { + REQUIRE_THROWS(ccf::kv::RawReader(nullptr, 1)); + } +} + +static std::vector make_public_domain_entry( + size_t declared_public_domain_size, const std::vector& public_domain) +{ + ccf::kv::SerialisedEntryHeader header; + header.set_size(sizeof(size_t) + public_domain.size()); + std::vector entry(sizeof(header) + header.size); + auto* data = entry.data(); + auto remaining = entry.size(); + serialized::write(data, remaining, header); + serialized::write(data, remaining, declared_public_domain_size); + serialized::write( + data, remaining, public_domain.data(), public_domain.size()); + return entry; +} + +TEST_CASE( + "KV deserialiser rejects invalid public domains" * + doctest::test_suite("serialisation")) +{ + auto encryptor = std::make_shared(); + + const auto initialise = [&](const std::vector& entry) { + ccf::kv::RawKvStoreDeserialiser deserialiser( + encryptor, ccf::kv::SecurityDomain::PUBLIC); + ccf::kv::Term term = 0; + ccf::kv::EntryFlags flags = {}; + return deserialiser.init(entry.data(), entry.size(), term, flags, false); + }; + + SUBCASE("Public domain exceeds remaining entry") + { + const auto entry = make_public_domain_entry(2, {0}); + REQUIRE_THROWS(initialise(entry)); + } + + SUBCASE("Public domain length prefix is truncated") + { + ccf::kv::SerialisedEntryHeader header; + header.set_size(sizeof(size_t) - 1); + std::vector entry(sizeof(header) + header.size); + std::memcpy(entry.data(), &header, sizeof(header)); + REQUIRE_THROWS(initialise(entry)); + } + + SUBCASE("Public domain length is impossible") + { + const auto entry = + make_public_domain_entry(std::numeric_limits::max(), {}); + REQUIRE_THROWS(initialise(entry)); + } + + SUBCASE("Claims digest is truncated") + { + std::vector public_domain( + sizeof(ccf::kv::EntryType) + sizeof(ccf::kv::Version)); + auto* data = public_domain.data(); + *data++ = static_cast(ccf::kv::EntryType::WriteSetWithClaims); + const ccf::kv::Version version = 1; + std::memcpy(data, &version, sizeof(version)); + const auto entry = + make_public_domain_entry(public_domain.size(), public_domain); + REQUIRE_THROWS(initialise(entry)); + } +} + TEST_CASE( "Serialise/deserialise public map only" * doctest::test_suite("serialisation"))