From 70505c331d12551cce268a1417c90b8a698222d3 Mon Sep 17 00:00:00 2001 From: Lee Rhodes Date: Thu, 24 Sep 2026 10:28:02 -0700 Subject: [PATCH] Validate header fields and sizes when deserializing CPC, Count-Min, Theta and VarOpt - var_opt_union: require the full non-empty preamble before reading it. - count_min_sketch: include the preamble in the deserialized size check; compute num_buckets * num_hashes in 64 bits. - compact theta, serial version 4: compute the packed data size in 64 bits, and require entry bits in [1, 63] and num entries bytes in [1, 4], for both byte and stream deserialization. - cpc_sketch: validate lg_k, num_coupons and the number of table entries; bound decoding by the number of compressed words; require decoded rows and columns to be in range; check the input size before allocating. Co-Authored-By: Claude Opus 5.5 --- count/include/count_min_impl.hpp | 7 ++-- count/test/count_min_test.cpp | 20 ++++++++++ cpc/include/cpc_compressor_impl.hpp | 21 +++++++--- cpc/include/cpc_sketch.hpp | 1 + cpc/include/cpc_sketch_impl.hpp | 21 +++++++++- cpc/test/cpc_sketch_test.cpp | 35 +++++++++++++++++ sampling/include/var_opt_union_impl.hpp | 1 + sampling/test/var_opt_union_test.cpp | 12 ++++++ theta/include/compact_theta_sketch_parser.hpp | 2 + .../compact_theta_sketch_parser_impl.hpp | 21 +++++++++- theta/include/theta_sketch_impl.hpp | 2 + theta/test/theta_sketch_test.cpp | 38 +++++++++++++++++++ 12 files changed, 168 insertions(+), 13 deletions(-) diff --git a/count/include/count_min_impl.hpp b/count/include/count_min_impl.hpp index f00e6457..dd6e8730 100644 --- a/count/include/count_min_impl.hpp +++ b/count/include/count_min_impl.hpp @@ -36,7 +36,7 @@ count_min_sketch::count_min_sketch(uint8_t num_hashes, uint32_t num_buckets _allocator(allocator), _num_hashes(num_hashes), _num_buckets(num_buckets), -_sketch_array((num_hashes*num_buckets < 1<<30) ? num_hashes*num_buckets : 0, 0, _allocator), +_sketch_array((static_cast(num_hashes) * num_buckets < 1<<30) ? static_cast(num_hashes) * num_buckets : 0, 0, _allocator), _seed(seed), _total_weight(0) { if (num_buckets < 3) { @@ -45,7 +45,7 @@ _total_weight(0) { // This check is to ensure later compatibility with a Java implementation whose maximum size can only // be 2^31-1. We check only against 2^30 for simplicity. - if (num_buckets * num_hashes >= 1 << 30) { + if (static_cast(num_buckets) * num_hashes >= 1 << 30) { throw std::invalid_argument("These parameters generate a sketch that exceeds 2^30 elements." "Try reducing either the number of buckets or the number of hash functions."); } @@ -407,7 +407,8 @@ auto count_min_sketch::deserialize(const void* bytes, size_t size, uint64_t const bool is_empty = (flags_byte & (1 << flags::IS_EMPTY)) > 0; if (is_empty) { return c; } // sketch is empty, no need to read further. - ensure_minimum_memory(size, sizeof(W) * (1 + nbuckets * nhashes)); + // preamble (already read) + total weight + table; nbuckets * nhashes < 2^30 was checked by the constructor + ensure_minimum_memory(size, PREAMBLE_LONGS_SHORT * sizeof(uint64_t) + sizeof(W) * (1 + c._sketch_array.size())); // Long 2 is the weight. W weight; diff --git a/count/test/count_min_test.cpp b/count/test/count_min_test.cpp index c884eb78..3e811f9c 100644 --- a/count/test/count_min_test.cpp +++ b/count/test/count_min_test.cpp @@ -336,4 +336,24 @@ TEST_CASE("CountMin sketch: sink serialize-deserialize round trip", "[cm_sketch] check_sink_serialize(non_empty); } +TEST_CASE("CountMin sketch: bytes deserialize truncated non-empty", "[cm_sketch]") { + count_min_sketch c(5, 64); + for (uint64_t i = 0; i < 10; ++i) c.update(i, 10 * i * i); + auto bytes = c.serialize(); + for (size_t trim = 1; trim <= 24; ++trim) { + REQUIRE_THROWS_AS(count_min_sketch::deserialize(bytes.data(), bytes.size() - trim), std::out_of_range); + } +} + +TEST_CASE("CountMin sketch: deserialize rejects overflowing dimensions", "[cm_sketch]") { + // num_buckets * num_hashes wraps to 0 in 32-bit arithmetic + count_min_sketch c(2, 64); + c.update(uint64_t(1)); + auto bytes = c.serialize(); + const uint32_t num_buckets = 1U << 31; + std::memcpy(bytes.data() + 8, &num_buckets, sizeof(num_buckets)); + REQUIRE_THROWS_AS(count_min_sketch::deserialize(bytes.data(), bytes.size()), std::invalid_argument); + REQUIRE_THROWS_AS(count_min_sketch(2, num_buckets), std::invalid_argument); +} + } /* namespace datasketches */ diff --git a/cpc/include/cpc_compressor_impl.hpp b/cpc/include/cpc_compressor_impl.hpp index 0cc24b19..601f5273 100644 --- a/cpc/include/cpc_compressor_impl.hpp +++ b/cpc/include/cpc_compressor_impl.hpp @@ -354,6 +354,7 @@ void cpc_compressor::uncompress_sliding_flavor(const compressed_state& sou const uint32_t row_col = pairs[i]; const uint32_t row = row_col >> 6; uint8_t col = row_col & 63; + if (col >= 56) throw std::out_of_range("col out of range"); // first undo the permutation col = permutation[col]; // then undo the rotation: old = (new + (offset+8)) mod 64 @@ -390,6 +391,9 @@ auto cpc_compressor::uncompress_surprising_values(const uint32_t* data, uint3 vector_u32 pairs(num_pairs, 0, allocator); const uint8_t num_base_bits = golomb_choose_number_of_base_bits(k + num_pairs, num_pairs); low_level_uncompress_pairs(pairs.data(), num_pairs, num_base_bits, data, data_words); + for (uint32_t i = 0; i < num_pairs; i++) { + if ((pairs[i] >> 6) >= k) throw std::out_of_range("row index out of range"); + } return pairs; } @@ -472,8 +476,10 @@ static inline void maybe_flush_bitbuf(uint64_t& bitbuf, uint8_t& bufbits, uint32 } } -static inline void maybe_fill_bitbuf(uint64_t& bitbuf, uint8_t& bufbits, const uint32_t* wordarr, uint32_t& wordindex, uint8_t minbits) { +static inline void maybe_fill_bitbuf(uint64_t& bitbuf, uint8_t& bufbits, const uint32_t* wordarr, uint32_t& wordindex, + uint32_t numwords, uint8_t minbits) { if (bufbits < minbits) { + if (wordindex >= numwords) throw std::out_of_range("compressed data over-run"); bitbuf |= static_cast(wordarr[wordindex++]) << bufbits; bufbits += 32; } @@ -530,7 +536,7 @@ void cpc_compressor::low_level_uncompress_bytes( if (compressed_words == nullptr) throw std::logic_error("compressed_words == NULL"); for (uint32_t byte_index = 0; byte_index < num_bytes_to_decode; byte_index++) { - maybe_fill_bitbuf(bitbuf, bufbits, compressed_words, word_index, 12); // ensure 12 bits in bit buffer + maybe_fill_bitbuf(bitbuf, bufbits, compressed_words, word_index, num_compressed_words, 12); // ensure 12 bits in bit buffer const size_t peek12 = bitbuf & 0xfff; // These 12 bits will include an entire Huffman codeword. const uint16_t lookup = decoding_table[peek12]; @@ -547,6 +553,7 @@ void cpc_compressor::low_level_uncompress_bytes( static inline uint64_t read_unary( const uint32_t* compressed_words, + uint32_t num_compressed_words, uint32_t& next_word_index, uint64_t& bitbuf, uint8_t& bufbits @@ -646,7 +653,7 @@ void cpc_compressor::low_level_uncompress_pairs( // y_delta_lo (basebits) for (uint32_t pair_index = 0; pair_index < num_pairs_to_decode; pair_index++) { - maybe_fill_bitbuf(bitbuf, bufbits, compressed_words, word_index, 12); // ensure 12 bits in bit buffer + maybe_fill_bitbuf(bitbuf, bufbits, compressed_words, word_index, num_compressed_words, 12); // ensure 12 bits in bit buffer const size_t peek12 = bitbuf & 0xfff; const uint16_t lookup = length_limited_unary_decoding_table65[peek12]; const uint8_t code_word_length = lookup >> 8; @@ -654,9 +661,9 @@ void cpc_compressor::low_level_uncompress_pairs( bitbuf >>= code_word_length; bufbits -= code_word_length; - const uint64_t golomb_hi = read_unary(compressed_words, word_index, bitbuf, bufbits); + const uint64_t golomb_hi = read_unary(compressed_words, num_compressed_words, word_index, bitbuf, bufbits); - maybe_fill_bitbuf(bitbuf, bufbits, compressed_words, word_index, num_base_bits); // ensure num_base_bits in bit buffer + maybe_fill_bitbuf(bitbuf, bufbits, compressed_words, word_index, num_compressed_words, num_base_bits); // ensure num_base_bits in bit buffer const uint64_t golomb_lo = bitbuf & golomb_lo_mask; bitbuf >>= num_base_bits; bufbits -= num_base_bits; @@ -666,6 +673,7 @@ void cpc_compressor::low_level_uncompress_pairs( if (y_delta > 0) predicted_col_index = 0; const uint32_t row_index = static_cast(predicted_row_index + y_delta); const uint8_t col_index = predicted_col_index + x_delta; + if (col_index > 63) throw std::out_of_range("column index out of range"); const uint32_t row_col = (row_index << 6) | col_index; pair_array[pair_index] = row_col; predicted_row_index = row_index; @@ -676,6 +684,7 @@ void cpc_compressor::low_level_uncompress_pairs( uint64_t read_unary( const uint32_t* compressed_words, + uint32_t num_compressed_words, uint32_t& next_word_index, uint64_t& bitbuf, uint8_t& bufbits @@ -683,7 +692,7 @@ uint64_t read_unary( if (compressed_words == nullptr) throw std::logic_error("compressed_words == NULL"); size_t subtotal = 0; while (true) { - maybe_fill_bitbuf(bitbuf, bufbits, compressed_words, next_word_index, 8); // ensure 8 bits in bit buffer + maybe_fill_bitbuf(bitbuf, bufbits, compressed_words, next_word_index, num_compressed_words, 8); // ensure 8 bits in bit buffer const uint8_t peek8 = bitbuf & 0xff; // These 8 bits include either all or part of the Unary codeword const uint8_t trailing_zeros = byte_trailing_zeros_table[peek8]; diff --git a/cpc/include/cpc_sketch.hpp b/cpc/include/cpc_sketch.hpp index b35e528d..9aba1f51 100644 --- a/cpc/include/cpc_sketch.hpp +++ b/cpc/include/cpc_sketch.hpp @@ -315,6 +315,7 @@ class cpc_sketch_alloc { inline size_t copy_hip_to_mem(void* dst) const; static void check_lg_k(uint8_t lg_k); + static void check_num_coupons(uint8_t lg_k, uint32_t num_coupons, uint32_t num_pairs); friend cpc_compressor; friend cpc_union_alloc; diff --git a/cpc/include/cpc_sketch_impl.hpp b/cpc/include/cpc_sketch_impl.hpp index 80f111f1..f805b121 100644 --- a/cpc/include/cpc_sketch_impl.hpp +++ b/cpc/include/cpc_sketch_impl.hpp @@ -527,6 +527,7 @@ cpc_sketch_alloc cpc_sketch_alloc::deserialize(std::istream& is, uint64_t const bool has_hip = flags_byte & (1 << flags::HAS_HIP); const bool has_table = flags_byte & (1 << flags::HAS_TABLE); const bool has_window = flags_byte & (1 << flags::HAS_WINDOW); + check_lg_k(lg_k); compressed_state compressed(allocator); compressed.table_data_words = 0; compressed.table_num_entries = 0; @@ -581,6 +582,7 @@ cpc_sketch_alloc cpc_sketch_alloc::deserialize(std::istream& is, uint64_t throw std::invalid_argument("Incompatible seed hashes: " + std::to_string(seed_hash) + ", " + std::to_string(compute_seed_hash(seed))); } + check_num_coupons(lg_k, num_coupons, compressed.table_num_entries); uncompressed_state uncompressed(allocator); get_compressor().uncompress(compressed, uncompressed, lg_k, num_coupons); if (!is.good()) { @@ -612,6 +614,7 @@ cpc_sketch_alloc cpc_sketch_alloc::deserialize(const void* bytes, size_t s const bool has_hip = flags_byte & (1 << flags::HAS_HIP); const bool has_table = flags_byte & (1 << flags::HAS_TABLE); const bool has_window = flags_byte & (1 << flags::HAS_WINDOW); + check_lg_k(lg_k); ensure_minimum_memory(size, preamble_ints << 2); compressed_state compressed(allocator); compressed.table_data_words = 0; @@ -646,13 +649,13 @@ cpc_sketch_alloc cpc_sketch_alloc::deserialize(const void* bytes, size_t s ptr += copy_from_mem(ptr, hip_est_accum); } if (has_window) { - compressed.window_data.resize(compressed.window_data_words); check_memory_size(ptr - base + (compressed.window_data_words * sizeof(uint32_t)), size); + compressed.window_data.resize(compressed.window_data_words); ptr += copy_from_mem(ptr, compressed.window_data.data(), compressed.window_data_words * sizeof(uint32_t)); } if (has_table) { - compressed.table_data.resize(compressed.table_data_words); check_memory_size(ptr - base + (compressed.table_data_words * sizeof(uint32_t)), size); + compressed.table_data.resize(compressed.table_data_words); ptr += copy_from_mem(ptr, compressed.table_data.data(), compressed.table_data_words * sizeof(uint32_t)); } if (!has_window) compressed.table_num_entries = num_coupons; @@ -676,6 +679,7 @@ cpc_sketch_alloc cpc_sketch_alloc::deserialize(const void* bytes, size_t s throw std::invalid_argument("Incompatible seed hashes: " + std::to_string(seed_hash) + ", " + std::to_string(compute_seed_hash(seed))); } + check_num_coupons(lg_k, num_coupons, compressed.table_num_entries); uncompressed_state uncompressed(allocator); get_compressor().uncompress(compressed, uncompressed, lg_k, num_coupons); return cpc_sketch_alloc(lg_k, num_coupons, first_interesting_column, std::move(uncompressed.table), @@ -720,6 +724,19 @@ size_t cpc_sketch_alloc::get_max_serialized_size_bytes(uint8_t lg_k) { return (int) (CPC_EMPIRICAL_MAX_SIZE_FACTOR * k) + CPC_MAX_PREAMBLE_SIZE_BYTES; } +template +void cpc_sketch_alloc::check_num_coupons(uint8_t lg_k, uint32_t num_coupons, uint32_t num_pairs) { + // at most one coupon per bit of the k x 64 bit matrix, and surprising values are a subset of coupons + if (num_coupons > (static_cast(1) << lg_k) * 64) { + throw std::invalid_argument("Possible corruption: num_coupons " + std::to_string(num_coupons) + + " exceeds the capacity for lg_k " + std::to_string(lg_k)); + } + if (num_pairs > num_coupons) { + throw std::invalid_argument("Possible corruption: table entries " + std::to_string(num_pairs) + + " exceed num_coupons " + std::to_string(num_coupons)); + } +} + template void cpc_sketch_alloc::check_lg_k(uint8_t lg_k) { if (lg_k < cpc_constants::MIN_LG_K || lg_k > cpc_constants::MAX_LG_K) { diff --git a/cpc/test/cpc_sketch_test.cpp b/cpc/test/cpc_sketch_test.cpp index e38d45cc..b46fbec6 100644 --- a/cpc/test/cpc_sketch_test.cpp +++ b/cpc/test/cpc_sketch_test.cpp @@ -379,4 +379,39 @@ TEST_CASE("cpc sketch: max serialized size", "[cpc_sketch]") { REQUIRE(cpc_sketch::get_max_serialized_size_bytes(26) == static_cast((0.6 * (1 << 26)) + 40)); } +TEST_CASE("cpc sketch: deserialize corrupt num_coupons", "[cpc_sketch]") { + cpc_sketch sketch(11); + for (int i = 0; i < 100; i++) sketch.update(i); + auto bytes = sketch.serialize(); + REQUIRE((bytes[5] & (1 << 3)) != 0); // sparse flavor: table present + REQUIRE((bytes[5] & (1 << 4)) == 0); // and no window + uint32_t num_coupons; + std::memcpy(&num_coupons, bytes.data() + 8, sizeof(num_coupons)); + + // more pairs than the compressed table holds: decoder must not read past it + auto corrupt = bytes; + const uint32_t more_coupons = num_coupons + 50; + std::memcpy(corrupt.data() + 8, &more_coupons, sizeof(more_coupons)); + REQUIRE_THROWS_AS(cpc_sketch::deserialize(corrupt.data(), corrupt.size()), std::out_of_range); + + // more coupons than the k x 64 bit matrix can hold + corrupt = bytes; + const uint32_t too_many_coupons = 64 * 2048 + 1; + std::memcpy(corrupt.data() + 8, &too_many_coupons, sizeof(too_many_coupons)); + REQUIRE_THROWS_AS(cpc_sketch::deserialize(corrupt.data(), corrupt.size()), std::invalid_argument); +} + +TEST_CASE("cpc sketch: deserialize corrupt lg_k", "[cpc_sketch]") { + cpc_sketch sketch(11); + for (int i = 0; i < 100; i++) sketch.update(i); + auto bytes = sketch.serialize(); + for (uint8_t lg_k: {0, 3, 27, 255}) { + bytes[3] = lg_k; + REQUIRE_THROWS_AS(cpc_sketch::deserialize(bytes.data(), bytes.size()), std::invalid_argument); + std::stringstream s; + s.write(reinterpret_cast(bytes.data()), bytes.size()); + REQUIRE_THROWS_AS(cpc_sketch::deserialize(s), std::invalid_argument); + } +} + } /* namespace datasketches */ diff --git a/sampling/include/var_opt_union_impl.hpp b/sampling/include/var_opt_union_impl.hpp index d04be0cb..d0529849 100644 --- a/sampling/include/var_opt_union_impl.hpp +++ b/sampling/include/var_opt_union_impl.hpp @@ -200,6 +200,7 @@ var_opt_union var_opt_union::deserialize(const void* bytes, size_t s return var_opt_union(max_k); } + ensure_minimum_memory(size, PREAMBLE_LONGS_NON_EMPTY << 3); uint64_t items_seen; ptr += copy_from_mem(ptr, items_seen); double outer_tau_numer; diff --git a/sampling/test/var_opt_union_test.cpp b/sampling/test/var_opt_union_test.cpp index b17d8fa4..a37a34f9 100644 --- a/sampling/test/var_opt_union_test.cpp +++ b/sampling/test/var_opt_union_test.cpp @@ -305,4 +305,16 @@ TEST_CASE("varopt union: serialize sampling", "[var_opt_union]") { compare_serialization_deserialization(u); } +TEST_CASE("varopt union: deserialize truncated non-empty preamble", "[var_opt_union]") { + var_opt_sketch sk(32); + sk.update(1); + var_opt_union u(32); + u.update(sk); + auto bytes = u.serialize(); + // a non-empty union needs 4 preamble longs + for (size_t size = 8; size < 32; ++size) { + REQUIRE_THROWS_AS(var_opt_union::deserialize(bytes.data(), size), std::out_of_range); + } +} + } diff --git a/theta/include/compact_theta_sketch_parser.hpp b/theta/include/compact_theta_sketch_parser.hpp index e5d9304e..d8bf3915 100644 --- a/theta/include/compact_theta_sketch_parser.hpp +++ b/theta/include/compact_theta_sketch_parser.hpp @@ -38,6 +38,8 @@ class compact_theta_sketch_parser { }; static compact_theta_sketch_data parse(const void* ptr, size_t size, uint64_t seed, bool dump_on_error = false); + static void check_v4_entry_bits(uint8_t entry_bits); + static void check_v4_num_entries_bytes(uint8_t num_entries_bytes); private: // offsets are in sizeof(type) diff --git a/theta/include/compact_theta_sketch_parser_impl.hpp b/theta/include/compact_theta_sketch_parser_impl.hpp index a801d1ea..1bddb7c9 100644 --- a/theta/include/compact_theta_sketch_parser_impl.hpp +++ b/theta/include/compact_theta_sketch_parser_impl.hpp @@ -49,6 +49,7 @@ auto compact_theta_sketch_parser::parse(const void* ptr, size_t size, uin theta = reinterpret_cast(ptr)[COMPACT_SKETCH_V4_THETA_U64]; } const uint8_t num_entries_bytes = reinterpret_cast(ptr)[COMPACT_SKETCH_V4_NUM_ENTRIES_BYTES_BYTE]; + check_v4_num_entries_bytes(num_entries_bytes); size_t data_offset_bytes = has_theta ? COMPACT_SKETCH_V4_PACKED_DATA_ESTIMATION_BYTE : COMPACT_SKETCH_V4_PACKED_DATA_EXACT_BYTE; check_memory_size(ptr, size, data_offset_bytes + num_entries_bytes, dump_on_error); uint32_t num_entries = 0; @@ -58,7 +59,8 @@ auto compact_theta_sketch_parser::parse(const void* ptr, size_t size, uin } data_offset_bytes += num_entries_bytes; const uint8_t entry_bits = reinterpret_cast(ptr)[COMPACT_SKETCH_V4_ENTRY_BITS_BYTE]; - const size_t expected_bits = entry_bits * num_entries; + check_v4_entry_bits(entry_bits); + const uint64_t expected_bits = static_cast(entry_bits) * num_entries; const size_t expected_size_bytes = data_offset_bytes + whole_bytes_to_hold_bits(expected_bits); check_memory_size(ptr, size, expected_size_bytes, dump_on_error); return {false, true, seed_hash, num_entries, theta, @@ -113,7 +115,7 @@ auto compact_theta_sketch_parser::parse(const void* ptr, size_t size, uin if (num_entries == 0) { return {true, true, seed_hash, 0, theta_constants::MAX_THETA, nullptr, 64}; } else { - const size_t expected_size_bytes = (preamble_size + num_entries) << 3; + const size_t expected_size_bytes = (preamble_size + static_cast(num_entries)) << 3; check_memory_size(ptr, size, expected_size_bytes, dump_on_error); const uint64_t* entries = reinterpret_cast(ptr) + COMPACT_SKETCH_ENTRIES_EXACT_U64; return {false, true, seed_hash, num_entries, theta_constants::MAX_THETA, entries, 64}; @@ -144,6 +146,21 @@ void compact_theta_sketch_parser::check_memory_size(const void* ptr, size + (dump_on_error ? (", sketch dump: " + hex_dump(reinterpret_cast(ptr), actual_bytes)) : "")); } +template +void compact_theta_sketch_parser::check_v4_entry_bits(uint8_t entry_bits) { + // deltas between ordered hashes below 2^63 need 1 to 63 bits + if (entry_bits == 0 || entry_bits > 63) { + throw std::invalid_argument("entry bits must be in [1, 63], actual " + std::to_string(entry_bits)); + } +} + +template +void compact_theta_sketch_parser::check_v4_num_entries_bytes(uint8_t num_entries_bytes) { + if (num_entries_bytes == 0 || num_entries_bytes > sizeof(uint32_t)) { + throw std::invalid_argument("num entries bytes must be in [1, 4], actual " + std::to_string(num_entries_bytes)); + } +} + template std::string compact_theta_sketch_parser::hex_dump(const uint8_t* ptr, size_t size) { std::stringstream s; diff --git a/theta/include/theta_sketch_impl.hpp b/theta/include/theta_sketch_impl.hpp index 0e341016..1ca7ea1d 100644 --- a/theta/include/theta_sketch_impl.hpp +++ b/theta/include/theta_sketch_impl.hpp @@ -696,7 +696,9 @@ compact_theta_sketch_alloc compact_theta_sketch_alloc::deserialize_v4( uint8_t preamble_longs, std::istream& is, uint64_t seed, const A& allocator) { const auto entry_bits = read(is); + compact_theta_sketch_parser::check_v4_entry_bits(entry_bits); const auto num_entries_bytes = read(is); + compact_theta_sketch_parser::check_v4_num_entries_bytes(num_entries_bytes); const auto flags_byte = read(is); const auto seed_hash = read(is); const bool is_empty = flags_byte & (1 << flags::IS_EMPTY); diff --git a/theta/test/theta_sketch_test.cpp b/theta/test/theta_sketch_test.cpp index adc0713f..29fca1f7 100644 --- a/theta/test/theta_sketch_test.cpp +++ b/theta/test/theta_sketch_test.cpp @@ -23,6 +23,7 @@ #include #include #include +#include #include #include @@ -902,4 +903,41 @@ TEST_CASE("max serialized size", "[theta_sketch]") { REQUIRE(max_size_bytes == compact_theta_sketch::get_max_serialized_size_bytes(lg_k)); } +TEST_CASE("theta sketch: deserialize v4 corrupt header", "[theta_sketch]") { + update_theta_sketch update_sketch = update_theta_sketch::builder().build(); + for (int i = 0; i < 100; ++i) update_sketch.update(i); + const auto bytes = update_sketch.compact().serialize_compressed(); + REQUIRE(bytes[1] == 4); + for (uint8_t entry_bits: {0, 64, 255}) { + auto corrupt = bytes; + corrupt[3] = entry_bits; + REQUIRE_THROWS_AS(compact_theta_sketch::deserialize(corrupt.data(), corrupt.size()), std::invalid_argument); + std::stringstream s; + s.write(reinterpret_cast(corrupt.data()), corrupt.size()); + REQUIRE_THROWS_AS(compact_theta_sketch::deserialize(s), std::invalid_argument); + } + for (uint8_t num_entries_bytes: {0, 5, 255}) { + auto corrupt = bytes; + corrupt[4] = num_entries_bytes; + REQUIRE_THROWS_AS(compact_theta_sketch::deserialize(corrupt.data(), corrupt.size()), std::invalid_argument); + std::stringstream s; + s.write(reinterpret_cast(corrupt.data()), corrupt.size()); + REQUIRE_THROWS_AS(compact_theta_sketch::deserialize(s), std::invalid_argument); + } +} + +TEST_CASE("theta sketch: deserialize v4 entry bits overflow", "[theta_sketch]") { + update_theta_sketch update_sketch = update_theta_sketch::builder().build(); + for (int i = 0; i < 100; ++i) update_sketch.update(i); + auto bytes = update_sketch.compact().serialize_compressed(); + REQUIRE(bytes[0] == 1); // exact mode, num_entries follows the first preamble long + // 8 bits * 2^29 entries = 2^32 bits wraps to 0 in 32-bit arithmetic + bytes.resize(12); + bytes[3] = 8; + bytes[4] = 4; + const uint32_t num_entries = 1U << 29; + std::memcpy(bytes.data() + 8, &num_entries, sizeof(num_entries)); + REQUIRE_THROWS_AS(compact_theta_sketch::deserialize(bytes.data(), bytes.size()), std::out_of_range); +} + } /* namespace datasketches */