Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions count/include/count_min_impl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@ count_min_sketch<W,A>::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<uint64_t>(num_hashes) * num_buckets < 1<<30) ? static_cast<size_t>(num_hashes) * num_buckets : 0, 0, _allocator),
_seed(seed),
_total_weight(0) {
if (num_buckets < 3) {
Expand All @@ -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<uint64_t>(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.");
}
Expand Down Expand Up @@ -407,7 +407,8 @@ auto count_min_sketch<W,A>::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;
Expand Down
20 changes: 20 additions & 0 deletions count/test/count_min_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<uint64_t> 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<uint64_t>::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<uint64_t> 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<uint64_t>::deserialize(bytes.data(), bytes.size()), std::invalid_argument);
REQUIRE_THROWS_AS(count_min_sketch<uint64_t>(2, num_buckets), std::invalid_argument);
}

} /* namespace datasketches */
21 changes: 15 additions & 6 deletions cpc/include/cpc_compressor_impl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -354,6 +354,7 @@ void cpc_compressor<A>::uncompress_sliding_flavor(const compressed_state<A>& 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
Expand Down Expand Up @@ -390,6 +391,9 @@ auto cpc_compressor<A>::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;
}

Expand Down Expand Up @@ -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<uint64_t>(wordarr[wordindex++]) << bufbits;
bufbits += 32;
}
Expand Down Expand Up @@ -530,7 +536,7 @@ void cpc_compressor<A>::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];
Expand All @@ -547,6 +553,7 @@ void cpc_compressor<A>::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
Expand Down Expand Up @@ -646,17 +653,17 @@ void cpc_compressor<A>::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;
const int8_t x_delta = lookup & 0xff;
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;
Expand All @@ -666,6 +673,7 @@ void cpc_compressor<A>::low_level_uncompress_pairs(
if (y_delta > 0) predicted_col_index = 0;
const uint32_t row_index = static_cast<uint32_t>(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;
Expand All @@ -676,14 +684,15 @@ void cpc_compressor<A>::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
) {
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];
Expand Down
1 change: 1 addition & 0 deletions cpc/include/cpc_sketch.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<A>;
friend cpc_union_alloc<A>;
Expand Down
21 changes: 19 additions & 2 deletions cpc/include/cpc_sketch_impl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -527,6 +527,7 @@ cpc_sketch_alloc<A> cpc_sketch_alloc<A>::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<A> compressed(allocator);
compressed.table_data_words = 0;
compressed.table_num_entries = 0;
Expand Down Expand Up @@ -581,6 +582,7 @@ cpc_sketch_alloc<A> cpc_sketch_alloc<A>::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<A> uncompressed(allocator);
get_compressor<A>().uncompress(compressed, uncompressed, lg_k, num_coupons);
if (!is.good()) {
Expand Down Expand Up @@ -612,6 +614,7 @@ cpc_sketch_alloc<A> cpc_sketch_alloc<A>::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<A> compressed(allocator);
compressed.table_data_words = 0;
Expand Down Expand Up @@ -646,13 +649,13 @@ cpc_sketch_alloc<A> cpc_sketch_alloc<A>::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;
Expand All @@ -676,6 +679,7 @@ cpc_sketch_alloc<A> cpc_sketch_alloc<A>::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<A> uncompressed(allocator);
get_compressor<A>().uncompress(compressed, uncompressed, lg_k, num_coupons);
return cpc_sketch_alloc(lg_k, num_coupons, first_interesting_column, std::move(uncompressed.table),
Expand Down Expand Up @@ -720,6 +724,19 @@ size_t cpc_sketch_alloc<A>::get_max_serialized_size_bytes(uint8_t lg_k) {
return (int) (CPC_EMPIRICAL_MAX_SIZE_FACTOR * k) + CPC_MAX_PREAMBLE_SIZE_BYTES;
}

template<typename A>
void cpc_sketch_alloc<A>::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<uint64_t>(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<typename A>
void cpc_sketch_alloc<A>::check_lg_k(uint8_t lg_k) {
if (lg_k < cpc_constants::MIN_LG_K || lg_k > cpc_constants::MAX_LG_K) {
Expand Down
35 changes: 35 additions & 0 deletions cpc/test/cpc_sketch_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -379,4 +379,39 @@ TEST_CASE("cpc sketch: max serialized size", "[cpc_sketch]") {
REQUIRE(cpc_sketch::get_max_serialized_size_bytes(26) == static_cast<size_t>((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<const char*>(bytes.data()), bytes.size());
REQUIRE_THROWS_AS(cpc_sketch::deserialize(s), std::invalid_argument);
}
}

} /* namespace datasketches */
1 change: 1 addition & 0 deletions sampling/include/var_opt_union_impl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -200,6 +200,7 @@ var_opt_union<T, A> var_opt_union<T, A>::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;
Expand Down
12 changes: 12 additions & 0 deletions sampling/test/var_opt_union_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<int> sk(32);
sk.update(1);
var_opt_union<int> 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<int>::deserialize(bytes.data(), size), std::out_of_range);
}
}

}
2 changes: 2 additions & 0 deletions theta/include/compact_theta_sketch_parser.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
21 changes: 19 additions & 2 deletions theta/include/compact_theta_sketch_parser_impl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ auto compact_theta_sketch_parser<dummy>::parse(const void* ptr, size_t size, uin
theta = reinterpret_cast<const uint64_t*>(ptr)[COMPACT_SKETCH_V4_THETA_U64];
}
const uint8_t num_entries_bytes = reinterpret_cast<const uint8_t*>(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;
Expand All @@ -58,7 +59,8 @@ auto compact_theta_sketch_parser<dummy>::parse(const void* ptr, size_t size, uin
}
data_offset_bytes += num_entries_bytes;
const uint8_t entry_bits = reinterpret_cast<const uint8_t*>(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<uint64_t>(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,
Expand Down Expand Up @@ -113,7 +115,7 @@ auto compact_theta_sketch_parser<dummy>::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<size_t>(num_entries)) << 3;
check_memory_size(ptr, size, expected_size_bytes, dump_on_error);
const uint64_t* entries = reinterpret_cast<const uint64_t*>(ptr) + COMPACT_SKETCH_ENTRIES_EXACT_U64;
return {false, true, seed_hash, num_entries, theta_constants::MAX_THETA, entries, 64};
Expand Down Expand Up @@ -144,6 +146,21 @@ void compact_theta_sketch_parser<dummy>::check_memory_size(const void* ptr, size
+ (dump_on_error ? (", sketch dump: " + hex_dump(reinterpret_cast<const uint8_t*>(ptr), actual_bytes)) : ""));
}

template<bool dummy>
void compact_theta_sketch_parser<dummy>::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<bool dummy>
void compact_theta_sketch_parser<dummy>::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<bool dummy>
std::string compact_theta_sketch_parser<dummy>::hex_dump(const uint8_t* ptr, size_t size) {
std::stringstream s;
Expand Down
2 changes: 2 additions & 0 deletions theta/include/theta_sketch_impl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -696,7 +696,9 @@ compact_theta_sketch_alloc<A> compact_theta_sketch_alloc<A>::deserialize_v4(
uint8_t preamble_longs, std::istream& is, uint64_t seed, const A& allocator)
{
const auto entry_bits = read<uint8_t>(is);
compact_theta_sketch_parser<true>::check_v4_entry_bits(entry_bits);
const auto num_entries_bytes = read<uint8_t>(is);
compact_theta_sketch_parser<true>::check_v4_num_entries_bytes(num_entries_bytes);
const auto flags_byte = read<uint8_t>(is);
const auto seed_hash = read<uint16_t>(is);
const bool is_empty = flags_byte & (1 << flags::IS_EMPTY);
Expand Down
Loading
Loading