diff --git a/merklecpp.h b/merklecpp.h index 12eba68..3a15ebb 100644 --- a/merklecpp.h +++ b/merklecpp.h @@ -1557,6 +1557,11 @@ namespace merkle const size_t deserialised_num_flushed = deserialise_size_t(bytes, position); + if (num_leaf_nodes == 0 && deserialised_num_flushed != 0) + { + throw std::runtime_error("serialised tree has no retained leaves"); + } + // A binary tree has 2 * leaves - 1 nodes, which must fit in Node::size. constexpr size_t max_num_leaves = std::numeric_limits::max() / 2 + 1; @@ -1579,17 +1584,19 @@ namespace merkle throw std::runtime_error("not enough bytes"); } - num_flushed = deserialised_num_flushed; - leaf_nodes.reserve(num_leaf_nodes); + std::vector deserialised_leaf_nodes; + deserialised_leaf_nodes.reserve(num_leaf_nodes); + std::vector> level; + level.reserve(num_leaf_nodes); for (size_t i = 0; i < num_leaf_nodes; i++) { - Node* n = Node::make(Hash(bytes, position)); - leaf_nodes.push_back(n); + auto n = std::unique_ptr(Node::make(Hash(bytes, position))); + deserialised_leaf_nodes.push_back(n.get()); + level.push_back(std::move(n)); } - std::vector level = leaf_nodes; - std::vector next_level; - size_t it = num_flushed; + std::vector> next_level; + size_t it = deserialised_num_flushed; uint8_t level_no = 0; while (it != 0 || level.size() > 1) { @@ -1598,11 +1605,11 @@ namespace merkle { Hash h(bytes, position); MERKLECPP_TRACE(MERKLECPP_TOUT << "+";); - auto n = Node::make(h); + auto n = std::unique_ptr(Node::make(std::move(h))); n->height = level_no + 1; n->size = Node::full_size(n->height); assert(n->invariant()); - level.insert(level.begin(), n); + level.insert(level.begin(), std::move(n)); } MERKLECPP_TRACE( @@ -1615,11 +1622,15 @@ namespace merkle { if (i + 1 >= level.size()) { - next_level.push_back(level.at(i)); + next_level.push_back(std::move(level.at(i))); } else { - next_level.push_back(Node::make(level.at(i), level.at(i + 1))); + auto parent = std::unique_ptr( + Node::make(level.at(i).get(), level.at(i + 1).get())); + level.at(i).release(); + level.at(i + 1).release(); + next_level.push_back(std::move(parent)); } } @@ -1634,9 +1645,11 @@ namespace merkle if (level.size() == 1) { - _root = level.at(0); + _root = level.at(0).release(); assert(_root->invariant()); } + leaf_nodes = std::move(deserialised_leaf_nodes); + num_flushed = deserialised_num_flushed; } /// @brief Operator to serialise the tree diff --git a/test/unit_tests.cpp b/test/unit_tests.cpp index 4d55735..c162ac1 100644 --- a/test/unit_tests.cpp +++ b/test/unit_tests.cpp @@ -296,6 +296,16 @@ TEST_CASE("TreeT rejects invalid serialised leaf data") "not enough bytes", std::runtime_error); + std::vector no_retained_leaves; + merkle::serialise_uint64_t(0, no_retained_leaves); + merkle::serialise_uint64_t(1, no_retained_leaves); + no_retained_leaves.resize( + no_retained_leaves.size() + merkle::Hash::size_bytes); + REQUIRE_THROWS_WITH_AS( + (void)merkle::Tree(no_retained_leaves), + "serialised tree has no retained leaves", + std::runtime_error); + std::vector overflowing_leaf_count; merkle::serialise_uint64_t(1, overflowing_leaf_count); merkle::serialise_uint64_t( @@ -325,6 +335,16 @@ TEST_CASE("TreeT rejects invalid serialised leaf data") (void)merkle::Tree(truncated_extra_hash), "not enough bytes", std::runtime_error); + + std::vector truncated_flushed_hashes; + merkle::serialise_uint64_t(1, truncated_flushed_hashes); + merkle::serialise_uint64_t(3, truncated_flushed_hashes); + truncated_flushed_hashes.resize( + truncated_flushed_hashes.size() + 2 * merkle::Hash::size_bytes); + REQUIRE_THROWS_WITH_AS( + (void)merkle::Tree(truncated_flushed_hashes), + "not enough bytes", + std::runtime_error); } TEST_CASE("TreeT deserialises flushed counts beyond signed shift width")