This is an automated email from the ASF dual-hosted git repository.
tisonkun pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/datasketches-rust.git
The following commit(s) were added to refs/heads/main by this push:
new 1841bf1 fix: simplify and harden sketch deserialization (#227)
1841bf1 is described below
commit 1841bf1f9af259640d60036de21782eff33af285
Author: tison <[email protected]>
AuthorDate: Wed Aug 26 16:54:16 2026 +0800
fix: simplify and harden sketch deserialization (#227)
Co-authored-by: Jaideep Pyne <[email protected]>
---
CHANGELOG.md | 5 +
datasketches/src/bloom/sketch.rs | 42 ++--
datasketches/src/cpc/compression.rs | 340 +++++++++++++++------------
datasketches/src/cpc/mod.rs | 12 +-
datasketches/src/cpc/pair_table.rs | 21 +-
datasketches/src/cpc/sketch.rs | 56 ++++-
datasketches/src/hll/array4.rs | 78 +++++-
datasketches/src/hll/array6.rs | 28 ++-
datasketches/src/hll/array8.rs | 28 ++-
datasketches/src/hll/hash_set.rs | 24 +-
datasketches/src/hll/list.rs | 19 +-
datasketches/src/hll/serialization.rs | 4 +-
datasketches/src/hll/sketch.rs | 20 +-
datasketches/src/req/compactor.rs | 40 ++--
datasketches/src/req/serialization.rs | 84 ++-----
datasketches/src/req/sketch.rs | 8 -
datasketches/src/thetafamily/theta/sketch.rs | 34 ++-
datasketches/src/thetafamily/tuple/sketch.rs | 9 +
datasketches/tests/cpc_test/deserialize.rs | 115 +++++++++
datasketches/tests/cpc_test/main.rs | 1 +
datasketches/tests/serde_tests/bloom.rs | 28 ++-
datasketches/tests/serde_tests/hll.rs | 75 ++++++
datasketches/tests/serde_tests/req.rs | 39 ++-
datasketches/tests/serde_tests/theta.rs | 23 ++
datasketches/tests/serde_tests/tuple.rs | 13 +
25 files changed, 805 insertions(+), 341 deletions(-)
diff --git a/CHANGELOG.md b/CHANGELOG.md
index 91ae5f6..269fa4d 100644
--- a/CHANGELOG.md
+++ b/CHANGELOG.md
@@ -18,9 +18,14 @@ All significant changes to this project will be documented
in this file.
### Bug fixes
+* Bloom filter deserialization now validates clean cached bit counts against
the bit array, reconstructs dirty counts, and checks non-empty payload sizes
before allocating.
* Frequent-items map sizes are now limited consistently to the cross-language
maximum of `2^30`: `FrequentItemsSketch::new` rejects larger configurations,
and `deserialize` returns `InvalidData` for out-of-range or inconsistent header
fields instead of panicking on corrupt input. Empty images restore the minimum
backing map instead of allocating from `lg_cur_map_size`, and non-empty images
validate `active_items` against the declared map capacity and remaining payload
before preallocatin [...]
* T-Digest compression now handles `k = u16::MAX` without overflowing the
scale normalization input.
* T-Digest deserialization now validates declared payload lengths before
allocating. Updating a deserialized digest whose unmerged buffer already
exceeds the compression threshold now compresses it instead of allowing the
buffer to grow without bound.
+* HLL deserialization now preserves HLL4 registers from compact images and
validates mode capacities and payload sizes before shifting or allocating.
+* Theta deserialization now validates declared entry counts and compressed
widths against the remaining payload before allocating or unpacking them.
+* Tuple deserialization now validates the declared entry count against the
remaining hash payload before allocating.
+* CPC deserialization now rejects malformed or corrupt input with an error
instead of panicking. Previously, corrupt bytes could trigger an
index-out-of-bounds, a failed assertion, or an arithmetic overflow while
decompressing the stream; the deserializer now validates the header fields
against the sketch flavor and bounds-checks the decompressor, while leaving
valid sketches byte-for-byte compatible with the Java, C++, and Go
implementations.
## v0.4.0 (2026-08-18)
diff --git a/datasketches/src/bloom/sketch.rs b/datasketches/src/bloom/sketch.rs
index 6eb3a45..af510f0 100644
--- a/datasketches/src/bloom/sketch.rs
+++ b/datasketches/src/bloom/sketch.rs
@@ -466,36 +466,40 @@ impl BloomFilter {
}
let num_words = num_longs as usize;
+ if !is_empty {
+ let payload_bytes = num_words
+ .checked_add(1)
+ .and_then(|words| words.checked_mul(size_of::<u64>()))
+ .ok_or_else(|| Error::deserial("Bloom filter payload length
overflows"))?;
+ if payload_bytes > cursor.remaining().len() {
+ return Err(Error::insufficient_data(format!(
+ "Bloom filter payload requires {payload_bytes} bytes, got
{}",
+ cursor.remaining().len()
+ )));
+ }
+ }
let mut bit_array = vec![0u64; num_words].into_boxed_slice();
- let num_bits_set;
-
- if is_empty {
- num_bits_set = 0;
+ let num_bits_set = if is_empty {
+ 0
} else {
- let raw_num_bits_set = cursor
+ let serialized_num_bits_set = cursor
.read_u64_le()
.map_err(insufficient_data("num_bits_set"))?;
- let mut counted_bits_set = 0;
+ let mut count = 0;
for word in &mut bit_array {
*word = cursor
.read_u64_le()
.map_err(insufficient_data("bit_array"))?;
- counted_bits_set += word.count_ones() as u64;
+ count += word.count_ones() as u64;
}
-
- // Handle "dirty" state: u64::MAX (all bits set to 1) indicates
bits need recounting.
- if raw_num_bits_set == u64::MAX {
- num_bits_set = counted_bits_set;
- } else {
- if raw_num_bits_set != counted_bits_set {
- return Err(Error::deserial(format!(
- "invalid num_bits_set: expected {counted_bits_set},
got {raw_num_bits_set}",
- )));
- }
- num_bits_set = raw_num_bits_set;
+ if serialized_num_bits_set != u64::MAX && serialized_num_bits_set
!= count {
+ return Err(Error::deserial(format!(
+ "invalid num_bits_set: expected {count}, got
{serialized_num_bits_set}"
+ )));
}
- }
+ count
+ };
Ok(BloomFilter {
seed,
diff --git a/datasketches/src/cpc/compression.rs
b/datasketches/src/cpc/compression.rs
index 752ad9d..db07f8e 100644
--- a/datasketches/src/cpc/compression.rs
+++ b/datasketches/src/cpc/compression.rs
@@ -28,6 +28,7 @@ use
crate::cpc::compression_data::LENGTH_LIMITED_UNARY_ENCODING_TABLE65;
use crate::cpc::determine_correct_offset;
use crate::cpc::determine_flavor;
use crate::cpc::pair_table::PairTable;
+use crate::error::Error;
#[derive(Default)]
pub(super) struct CompressedState {
@@ -355,12 +356,12 @@ pub(super) struct UncompressedState {
}
impl CompressedState {
- pub fn uncompress(&self, lg_k: u8, num_coupons: u32) -> UncompressedState {
+ pub fn uncompress(&self, lg_k: u8, num_coupons: u32) ->
Result<UncompressedState, Error> {
match determine_flavor(lg_k, num_coupons) {
- Flavor::Empty => UncompressedState {
+ Flavor::Empty => Ok(UncompressedState {
table: PairTable::new(2, lg_k + 6),
window: vec![],
- },
+ }),
Flavor::Sparse => self.uncompress_sparse_flavor(lg_k),
Flavor::Hybrid => self.uncompress_hybrid_flavor(lg_k),
Flavor::Pinned => self.uncompress_pinned_flavor(lg_k, num_coupons),
@@ -368,33 +369,31 @@ impl CompressedState {
}
}
- fn uncompress_sparse_flavor(&self, lg_k: u8) -> UncompressedState {
+ fn uncompress_sparse_flavor(&self, lg_k: u8) -> Result<UncompressedState,
Error> {
debug_assert!(self.window_data.is_empty(), "window is not expected");
- debug_assert!(!self.table_data.is_empty(), "table is expected");
let pairs = uncompress_surprising_values(
&self.table_data,
self.table_data_words,
self.table_num_entries,
lg_k,
- );
+ )?;
- UncompressedState {
- table: PairTable::from_slots(lg_k, self.table_num_entries, pairs),
+ Ok(UncompressedState {
+ table: PairTable::from_slots(lg_k, self.table_num_entries, pairs)?,
window: vec![],
- }
+ })
}
- fn uncompress_hybrid_flavor(&self, lg_k: u8) -> UncompressedState {
+ fn uncompress_hybrid_flavor(&self, lg_k: u8) -> Result<UncompressedState,
Error> {
debug_assert!(self.window_data.is_empty(), "window is not expected");
- debug_assert!(!self.table_data.is_empty(), "table is expected");
let mut pairs = uncompress_surprising_values(
&self.table_data,
self.table_data_words,
self.table_num_entries,
lg_k,
- );
+ )?;
// In the hybrid flavor, some of these pairs actually belong in the
window, so we will
// separate them out, moving the "true" pairs to the bottom of the
array.
@@ -403,7 +402,6 @@ impl CompressedState {
let mut next_true_pair = 0;
for i in 0..self.table_num_entries {
let row_col = pairs[i as usize];
- assert_ne!(row_col, u32::MAX);
let col = row_col & 63;
if col < 8 {
let row = row_col >> 6;
@@ -414,15 +412,17 @@ impl CompressedState {
}
}
- UncompressedState {
- table: PairTable::from_slots(lg_k, next_true_pair, pairs),
+ Ok(UncompressedState {
+ table: PairTable::from_slots(lg_k, next_true_pair, pairs)?,
window,
- }
+ })
}
- fn uncompress_pinned_flavor(&self, lg_k: u8, num_coupons: u32) ->
UncompressedState {
- debug_assert!(!self.window_data.is_empty(), "window is expected");
-
+ fn uncompress_pinned_flavor(
+ &self,
+ lg_k: u8,
+ num_coupons: u32,
+ ) -> Result<UncompressedState, Error> {
let mut window = vec![];
uncompress_sliding_window(
&self.window_data,
@@ -430,36 +430,38 @@ impl CompressedState {
&mut window,
lg_k,
num_coupons,
- );
+ )?;
let num_pairs = self.table_num_entries;
let table = if num_pairs == 0 {
PairTable::new(2, lg_k + 6)
} else {
- debug_assert!(!self.table_data.is_empty(), "table is expected");
let mut pairs = uncompress_surprising_values(
&self.table_data,
self.table_data_words,
num_pairs,
lg_k,
- );
+ )?;
// undo the compressor's 8-column shift
for i in 0..num_pairs {
let i = i as usize;
- assert!(
- (pairs[i] & 63) < 56,
- "pair column index is invalid: {}",
- pairs[i]
- );
+ if (pairs[i] & 63) >= 56 {
+ return Err(Error::deserial(format!(
+ "CPC pinned table pair column index is invalid: {}",
+ pairs[i]
+ )));
+ }
pairs[i] += 8;
}
- PairTable::from_slots(lg_k, num_pairs, pairs)
+ PairTable::from_slots(lg_k, num_pairs, pairs)?
};
- UncompressedState { table, window }
+ Ok(UncompressedState { table, window })
}
- fn uncompress_sliding_flavor(&self, lg_k: u8, num_coupons: u32) ->
UncompressedState {
- debug_assert!(!self.window_data.is_empty(), "window is expected");
-
+ fn uncompress_sliding_flavor(
+ &self,
+ lg_k: u8,
+ num_coupons: u32,
+ ) -> Result<UncompressedState, Error> {
let mut window = vec![];
uncompress_sliding_window(
&self.window_data,
@@ -467,38 +469,46 @@ impl CompressedState {
&mut window,
lg_k,
num_coupons,
- );
+ )?;
let num_pairs = self.table_num_entries;
let table = if num_pairs == 0 {
PairTable::new(2, lg_k + 6)
} else {
- debug_assert!(!self.table_data.is_empty(), "table is expected");
let mut pairs = uncompress_surprising_values(
&self.table_data,
self.table_data_words,
num_pairs,
lg_k,
- );
+ )?;
let pseudo_phase = determine_pseudo_phase(lg_k, num_coupons);
let permutation = &COLUMN_PERMUTATIONS_FOR_DECODING[pseudo_phase
as usize];
let offset = determine_correct_offset(lg_k, num_coupons);
- assert!(offset <= 56, "offset is invalid: {offset}");
+ if offset > 56 {
+ return Err(Error::deserial(format!(
+ "CPC sliding window offset is invalid: {offset}"
+ )));
+ }
for i in 0..num_pairs {
let i = i as usize;
let row_col = pairs[i];
let row = row_col >> 6;
- let mut col = (row_col & 63) as u8;
+ let col = (row_col & 63) as usize;
// first undo the permutation
- col = permutation[col as usize];
+ if col >= permutation.len() {
+ return Err(Error::deserial(format!(
+ "CPC sliding table pair column index is invalid: {}",
+ pairs[i]
+ )));
+ }
+ let mut col = permutation[col];
// then undo the rotation: old = (new + (offset+8)) mod 64
col = (col + (offset + 8)) & 63;
pairs[i] = (row << 6) | (col as u32);
}
-
- PairTable::from_slots(lg_k, num_pairs, pairs)
+ PairTable::from_slots(lg_k, num_pairs, pairs)?
};
- UncompressedState { table, window }
+ Ok(UncompressedState { table, window })
}
}
@@ -507,12 +517,15 @@ fn uncompress_surprising_values(
data_words: usize,
num_pairs: u32,
lg_k: u8,
-) -> Vec<u32> {
+) -> Result<Vec<u32>, Error> {
+ if num_pairs == 0 {
+ return Ok(vec![]);
+ }
let k = 1 << lg_k;
let mut pairs = vec![0; num_pairs as usize];
let num_base_bits = golomb_choose_number_of_base_bits(k + num_pairs,
num_pairs as u64);
- low_level_uncompress_pairs(&mut pairs, num_pairs, num_base_bits, data,
data_words);
- pairs
+ low_level_uncompress_pairs(&mut pairs, num_pairs, k, num_base_bits, data,
data_words)?;
+ Ok(pairs)
}
fn uncompress_sliding_window(
@@ -521,7 +534,7 @@ fn uncompress_sliding_window(
window: &mut Vec<u8>,
lg_k: u8,
num_coupons: u32,
-) {
+) -> Result<(), Error> {
let k = 1 << lg_k;
window.resize(k, 0);
let pseudo_phase = determine_pseudo_phase(lg_k, num_coupons);
@@ -531,22 +544,21 @@ fn uncompress_sliding_window(
data,
data_words,
&DECODING_TABLES_FOR_HIGH_ENTROPY_BYTE[pseudo_phase as usize],
- );
+ )
}
fn low_level_uncompress_pairs(
pairs: &mut [u32],
num_pairs_to_decode: u32,
+ k: u32,
num_base_bits: u8,
compressed_words: &[u32],
num_compressed_words: usize,
-) {
- let mut word_index = 0;
- let mut bitbuf = 0;
- let mut bufbits = 0;
+) -> Result<(), Error> {
+ let mut bits = BitReader::new(compressed_words, num_compressed_words)?;
let golomb_lo_mask = (1 << num_base_bits) - 1;
let mut predicted_row_index = 0u32;
- let mut predicted_col_index = 0u8;
+ let mut predicted_col_index = 0u32;
// for each pair we need to read:
// x_delta (12-bit length-limited unary)
@@ -554,51 +566,43 @@ fn low_level_uncompress_pairs(
// y_delta_lo (basebits)
for pair_index in 0..num_pairs_to_decode {
- // ensure 12 bits in bit buffer
- maybe_fill_bitbuf(
- &mut bitbuf,
- &mut bufbits,
- compressed_words,
- &mut word_index,
- 12,
- );
- let peek12 = bitbuf & 0xfff;
+ let peek12 = bits.peek(12)?;
let lookup = LENGTH_LIMITED_UNARY_DECODING_TABLE65[peek12 as usize];
let code_word_length = (lookup >> 8) as u8;
- let x_delta = (lookup & 0xff) as u8;
- bitbuf >>= code_word_length;
- bufbits -= code_word_length;
+ let x_delta = u32::from((lookup & 0xff) as u8);
+ bits.consume(code_word_length);
- let golomb_hi = read_unary(compressed_words, &mut word_index, &mut
bitbuf, &mut bufbits);
- // ensure num_base_bits in the bit buffer
- maybe_fill_bitbuf(
- &mut bitbuf,
- &mut bufbits,
- compressed_words,
- &mut word_index,
- num_base_bits,
- );
- let golomb_lo = bitbuf & golomb_lo_mask;
- bitbuf >>= num_base_bits;
- bufbits -= num_base_bits;
- let y_delta = ((golomb_hi << num_base_bits) | golomb_lo) as u32;
+ let golomb_hi = bits.read_unary()?;
+ let golomb_lo = bits.read(num_base_bits)? & golomb_lo_mask;
+ let y_delta = golomb_hi
+ .checked_shl(u32::from(num_base_bits))
+ .and_then(|high| high.checked_add(golomb_lo))
+ .and_then(|delta| u32::try_from(delta).ok())
+ .ok_or_else(|| Error::deserial("CPC pair row delta overflows"))?;
// Now that we have x_delta and y_delta, we can compute the pair's row
and column
if y_delta > 0 {
predicted_col_index = 0;
}
- let row_index = predicted_row_index + y_delta;
- let col_index = predicted_col_index + x_delta;
- let row_col = (row_index << 6) | (col_index as u32);
+ let row_index = predicted_row_index
+ .checked_add(y_delta)
+ .filter(|&row| row < k)
+ .ok_or_else(|| Error::deserial("CPC pair row index is out of
range"))?;
+ let col_index = predicted_col_index
+ .checked_add(x_delta)
+ .filter(|&column| column < 64)
+ .ok_or_else(|| Error::deserial("CPC pair column index is out of
range"))?;
+ let row_col = (row_index << 6) | col_index;
+ if row_col == u32::MAX {
+ return Err(Error::deserial(
+ "CPC pair uses the reserved empty-table sentinel",
+ ));
+ }
pairs[pair_index as usize] = row_col;
predicted_row_index = row_index;
predicted_col_index = col_index + 1;
}
-
- debug_assert!(
- word_index <= num_compressed_words,
- "word_index: {word_index}, num_compressed_words:
{num_compressed_words}",
- );
+ Ok(())
}
fn low_level_uncompress_bytes(
@@ -607,39 +611,96 @@ fn low_level_uncompress_bytes(
compressed_words: &[u32],
num_compressed_words: usize,
decoding_table: &[u16],
-) {
- let mut word_index = 0;
- let mut bitbuf = 0;
- let mut bufbits = 0;
+) -> Result<(), Error> {
+ let mut bits = BitReader::new(compressed_words, num_compressed_words)?;
for byte_index in 0..num_bytes_to_decode {
- // ensure 12 bits in bit buffer
- maybe_fill_bitbuf(
- &mut bitbuf,
- &mut bufbits,
- compressed_words,
- &mut word_index,
- 12,
- );
// These 12 bits will include an entire Huffman codeword.
- let peek12 = bitbuf & 0xfff;
+ let peek12 = bits.peek(12)?;
let lookup = decoding_table[peek12 as usize];
let code_word_length = (lookup >> 8) as u8;
let decoded_byte = (lookup & 0xff) as u8;
byte_array[byte_index as usize] = decoded_byte;
- bitbuf >>= code_word_length;
- bufbits -= code_word_length;
+ bits.consume(code_word_length);
+ }
+ Ok(())
+}
+
+struct BitReader<'a> {
+ words: &'a [u32],
+ next_word: usize,
+ buffer: u64,
+ buffered_bits: u8,
+}
+
+impl<'a> BitReader<'a> {
+ fn new(words: &'a [u32], num_words: usize) -> Result<Self, Error> {
+ let words = words
+ .get(..num_words)
+ .ok_or_else(|| Error::deserial("CPC compressed word count exceeds
payload"))?;
+ Ok(Self {
+ words,
+ next_word: 0,
+ buffer: 0,
+ buffered_bits: 0,
+ })
+ }
+
+ fn fill(&mut self, needed: u8) -> Result<(), Error> {
+ if self.buffered_bits < needed {
+ let word = self
+ .words
+ .get(self.next_word)
+ .ok_or_else(|| Error::deserial("CPC compressed stream is
truncated"))?;
+ self.buffer |= u64::from(*word) << self.buffered_bits;
+ self.next_word += 1;
+ self.buffered_bits += 32;
+ }
+ Ok(())
+ }
+
+ fn peek(&mut self, count: u8) -> Result<u64, Error> {
+ self.fill(count)?;
+ Ok(self.buffer & ((1u64 << count) - 1))
+ }
+
+ fn consume(&mut self, count: u8) {
+ debug_assert!(count <= self.buffered_bits);
+ self.buffer >>= count;
+ self.buffered_bits -= count;
}
- // Buffer over-run should be impossible unless there is a bug.
- debug_assert!(
- word_index <= num_compressed_words,
- "word_index: {word_index}, num_compressed_words:
{num_compressed_words}",
- );
+ fn read(&mut self, count: u8) -> Result<u64, Error> {
+ if count == 0 {
+ return Ok(0);
+ }
+ let value = self.peek(count)?;
+ self.consume(count);
+ Ok(value)
+ }
+
+ fn read_unary(&mut self) -> Result<u64, Error> {
+ let mut value = 0u64;
+ loop {
+ let byte = self.peek(8)?;
+ let zeros = byte.trailing_zeros() as u8;
+ if zeros < 8 {
+ self.consume(zeros + 1);
+ return value
+ .checked_add(u64::from(zeros))
+ .ok_or_else(|| Error::deserial("CPC unary value
overflows"));
+ }
+ value = value
+ .checked_add(8)
+ .ok_or_else(|| Error::deserial("CPC unary value overflows"))?;
+ self.consume(8);
+ }
+ }
}
fn determine_pseudo_phase(lg_k: u8, num_coupons: u32) -> u8 {
- let k = 1 << lg_k;
+ let k = 1u64 << lg_k;
+ let num_coupons = u64::from(num_coupons);
// This mid-range logic produces pseudo-phases. They are used to select
encoding tables.
// The thresholds were chosen by hand after looking at plots of measured
compression.
if 1000 * num_coupons < 2375 * k {
@@ -698,31 +759,6 @@ fn write_unary(
maybe_flush_bitbuf(bitbuf, bufbits, compressed_words, next_word_index);
}
-fn read_unary(
- compressed_words: &[u32],
- next_word_index: &mut usize,
- bitbuf: &mut u64,
- bufbits: &mut u8,
-) -> u64 {
- let mut subtotal = 0u64;
- loop {
- // ensure 8 bits in bit buffer
- maybe_fill_bitbuf(bitbuf, bufbits, compressed_words, next_word_index,
8);
- // These 8 bits include either all or part of the Unary codeword
- let peek8 = *bitbuf & 0xff;
- let trailing_zeros = peek8.trailing_zeros() as u8;
- if trailing_zeros < 8 {
- *bufbits -= 1 + trailing_zeros;
- *bitbuf >>= 1 + trailing_zeros;
- return subtotal + trailing_zeros as u64;
- }
- // The codeword was partial, so read some more
- subtotal += 8;
- *bufbits -= 8;
- *bitbuf >>= 8;
- }
-}
-
fn maybe_flush_bitbuf(
bitbuf: &mut u64,
bufbits: &mut u8,
@@ -737,20 +773,6 @@ fn maybe_flush_bitbuf(
}
}
-fn maybe_fill_bitbuf(
- bitbuf: &mut u64,
- bufbits: &mut u8,
- words: &[u32],
- word_index: &mut usize,
- minbits: u8,
-) {
- if *bufbits < minbits {
- *bitbuf |= (words[*word_index] as u64) << *bufbits;
- *word_index += 1;
- *bufbits += 32;
- }
-}
-
// Explanation of padding: we write
// 1) xdelta (huffman, provides at least 1 bit, requires 12-bit lookahead)
// 2) ydeltaGolombHi (unary, provides at least 1 bit, requires 8-bit lookahead)
@@ -816,3 +838,33 @@ fn floor_log2_of_long(x: u64) -> u8 {
}
}
}
+
+#[cfg(test)]
+mod tests {
+ use super::CompressedState;
+ use super::determine_pseudo_phase;
+ use super::uncompress_surprising_values;
+
+ #[test]
+ fn pseudo_phase_handles_maximum_lg_k() {
+ assert!(determine_pseudo_phase(26, 1 << 25) < 22);
+ assert!(determine_pseudo_phase(26, u32::MAX) < 22);
+ }
+
+ #[test]
+ fn pair_decoder_rejects_empty_table_sentinel() {
+ let mut compressed = CompressedState::default();
+ compressed.compress_surprising_values(&[u32::MAX], 26);
+
+ let error = uncompress_surprising_values(
+ &compressed.table_data,
+ compressed.table_data_words,
+ 1,
+ 26,
+ );
+ assert_eq!(
+ error.unwrap_err().message(),
+ "CPC pair uses the reserved empty-table sentinel"
+ );
+ }
+}
diff --git a/datasketches/src/cpc/mod.rs b/datasketches/src/cpc/mod.rs
index 0313c98..6102dea 100644
--- a/datasketches/src/cpc/mod.rs
+++ b/datasketches/src/cpc/mod.rs
@@ -74,17 +74,15 @@ fn count_bits_set_in_matrix(matrix: &[u64]) -> u32 {
}
fn determine_flavor(lg_k: u8, num_coupons: u32) -> Flavor {
- let k = 1 << lg_k;
- let c2 = num_coupons << 1;
- let c8 = num_coupons << 3;
- let c32 = num_coupons << 5;
+ let k = 1u64 << lg_k;
+ let coupons = u64::from(num_coupons);
if num_coupons == 0 {
Flavor::Empty
- } else if c32 < (3 * k) {
+ } else if 32 * coupons < 3 * k {
Flavor::Sparse
- } else if c2 < k {
+ } else if 2 * coupons < k {
Flavor::Hybrid
- } else if c8 < (27 * k) {
+ } else if 8 * coupons < 27 * k {
Flavor::Pinned
} else {
Flavor::Sliding
diff --git a/datasketches/src/cpc/pair_table.rs
b/datasketches/src/cpc/pair_table.rs
index 4c4de42..8c0fb62 100644
--- a/datasketches/src/cpc/pair_table.rs
+++ b/datasketches/src/cpc/pair_table.rs
@@ -15,6 +15,8 @@
// specific language governing permissions and limitations
// under the License.
+use crate::error::Error;
+
const UPSIZE_NUMERATOR: u32 = 3;
const UPSIZE_DENOMINATOR: u32 = 4;
const DOWNSIZE_NUMERATOR: u32 = 1;
@@ -52,9 +54,17 @@ impl PairTable {
}
/// A constructor specifically tailored to be a part of FM85 decompression
scheme.
- pub fn from_slots(lg_size: u8, num_items: u32, slots: Vec<u32>) -> Self {
+ pub fn from_slots(lg_size: u8, num_items: u32, slots: Vec<u32>) ->
Result<Self, Error> {
+ if slots.len() < num_items as usize {
+ return Err(Error::deserial("CPC pair table contains too few
slots"));
+ }
let mut lg_num_slots = 2;
- while UPSIZE_DENOMINATOR * num_items > (UPSIZE_NUMERATOR * (1 <<
lg_num_slots)) {
+ while u64::from(UPSIZE_DENOMINATOR) * u64::from(num_items)
+ > u64::from(UPSIZE_NUMERATOR) * (1u64 << lg_num_slots)
+ {
+ if lg_num_slots == 26 {
+ return Err(Error::deserial("CPC pair table is too large"));
+ }
lg_num_slots += 1;
}
@@ -65,10 +75,11 @@ impl PairTable {
// the problem might not occur.
for i in 0..num_items {
- table.must_insert(slots[i as usize]);
+ if !table.maybe_insert(slots[i as usize]) {
+ return Err(Error::deserial("CPC pair table contains a
duplicate pair"));
+ }
}
- table.num_items = num_items;
- table
+ Ok(table)
}
pub fn slots(&self) -> &[u32] {
diff --git a/datasketches/src/cpc/sketch.rs b/datasketches/src/cpc/sketch.rs
index 622b6a1..2ae9164 100644
--- a/datasketches/src/cpc/sketch.rs
+++ b/datasketches/src/cpc/sketch.rs
@@ -639,7 +639,61 @@ impl CpcSketch {
)));
}
- let uncompressed = compressed.uncompress(lg_k, num_coupons);
+ // The coupon space of a sketch has `k * 64` cells (`k` rows of 64
columns each), so a
+ // valid sketch can never report more coupons than that. Rejecting
larger values keeps the
+ // flavor arithmetic below from overflowing on corrupt input.
+ if (num_coupons as u64) > 64 * (1u64 << lg_k) {
+ return Err(Error::deserial(format!(
+ "num_coupons ({}) exceeds coupon space for lg_k = {}",
+ num_coupons, lg_k
+ )));
+ }
+
+ // A valid sketch stores a sliding window exactly for the pinned and
sliding flavors, and
+ // stores its coupons in the surprising-value table for the sparse and
hybrid flavors. The
+ // flavor is fully determined by `lg_k` and `num_coupons`, so the
flags must agree with it.
+ let flavor = determine_flavor(lg_k, num_coupons);
+ let window_expected = matches!(flavor, Flavor::Pinned |
Flavor::Sliding);
+ if has_window != window_expected {
+ return Err(Error::deserial(format!(
+ "sliding-window flag ({}) is inconsistent with the {:?}
flavor",
+ has_window, flavor
+ )));
+ }
+ if matches!(flavor, Flavor::Sparse | Flavor::Hybrid) && !has_table {
+ return Err(Error::deserial(format!(
+ "table flag is unset but required for the {:?} flavor",
+ flavor
+ )));
+ }
+
+ // The number of stored table entries can never exceed the number of
coupons.
+ if compressed.table_num_entries > num_coupons {
+ return Err(Error::deserial(format!(
+ "table_num_entries ({}) exceeds num_coupons ({})",
+ compressed.table_num_entries, num_coupons
+ )));
+ }
+ // A pair requires at least one bit to encode, so the declared number
of table entries can
+ // never exceed the number of bits available in the table data. This
also bounds the size
+ // of the allocation made while decoding, rejecting corrupt inputs
that claim an enormous
+ // entry count backed by only a few data words.
+ if (compressed.table_num_entries as usize) >
compressed.table_data_words.saturating_mul(32)
+ {
+ return Err(Error::deserial(format!(
+ "table_num_entries ({}) exceeds capacity of table data ({}
words)",
+ compressed.table_num_entries, compressed.table_data_words
+ )));
+ }
+ let k = 1usize << lg_k;
+ if has_window && compressed.window_data_words.saturating_mul(32) < k {
+ return Err(Error::deserial(format!(
+ "window data ({} words) is too short for lg_k = {lg_k}",
+ compressed.window_data_words
+ )));
+ }
+
+ let uncompressed = compressed.uncompress(lg_k, num_coupons)?;
Ok(CpcSketch {
lg_k,
seed,
diff --git a/datasketches/src/hll/array4.rs b/datasketches/src/hll/array4.rs
index c2d0413..17327fb 100644
--- a/datasketches/src/hll/array4.rs
+++ b/datasketches/src/hll/array4.rs
@@ -41,6 +41,22 @@ use crate::hll::serialization::encode_mode_byte;
const AUX_TOKEN: u8 = 15;
+#[derive(Clone, Copy)]
+pub(super) enum AuxFormat {
+ Compact,
+ Updatable { lg_arr: u8 },
+}
+
+impl AuxFormat {
+ pub(super) fn from_header(compact: bool, lg_arr: u8) -> Self {
+ if compact {
+ Self::Compact
+ } else {
+ Self::Updatable { lg_arr }
+ }
+ }
+}
+
/// Core Array4 data structure - stores 4-bit values efficiently
#[derive(Debug, Clone, PartialEq)]
pub struct Array4 {
@@ -300,10 +316,11 @@ impl Array4 {
mut cursor: SketchSlice,
cur_min: u8,
lg_config_k: u8,
- compact: bool,
+ aux_format: AuxFormat,
ooo: bool,
) -> Result<Self, Error> {
- let num_bytes = 1 << (lg_config_k - 1); // k/2 bytes for 4-bit packing
+ let k = 1usize << lg_config_k;
+ let num_bytes = 1usize << (lg_config_k - 1); // k/2 bytes for 4-bit
packing
// Read HIP estimator values from preamble
let hip_accum = cursor
@@ -319,31 +336,70 @@ impl Array4 {
let aux_count = cursor
.read_u32_le()
.map_err(insufficient_data("aux_count"))?;
+ if num_at_cur_min as usize > k || aux_count as usize > k {
+ return Err(Error::deserial(
+ "HLL4 register or auxiliary count exceeds k",
+ ));
+ }
+ let (aux_slots, compact) = match aux_format {
+ AuxFormat::Compact => (aux_count as usize, true),
+ AuxFormat::Updatable { lg_arr } => {
+ let slots = if aux_count == 0 {
+ 0
+ } else {
+ 1usize
+ .checked_shl(u32::from(lg_arr))
+ .ok_or_else(|| Error::deserial(format!("invalid HLL4
lg_arr: {lg_arr}")))?
+ };
+ (slots, false)
+ }
+ };
+ let required_bytes = aux_slots
+ .checked_mul(COUPON_SIZE_BYTES)
+ .and_then(|aux_bytes| num_bytes.checked_add(aux_bytes))
+ .ok_or_else(|| Error::deserial("HLL4 payload length overflows"))?;
+ if required_bytes > cursor.remaining().len() {
+ return Err(Error::insufficient_data(format!(
+ "HLL4 payload requires {required_bytes} bytes, got {}",
+ cursor.remaining().len()
+ )));
+ }
// Read packed 4-bit byte array
let mut data = vec![0u8; num_bytes];
- if !compact {
- cursor
- .read_exact(&mut data)
- .map_err(insufficient_data("data"))?;
- } else {
- cursor.advance(num_bytes as u64);
- }
+ cursor
+ .read_exact(&mut data)
+ .map_err(insufficient_data("data"))?;
// Read aux map if present
let mut aux_map = None;
if aux_count > 0 {
let mut aux = AuxMap::new(lg_config_k);
- for i in 0..aux_count {
+ let mut decoded_count = 0;
+ for i in 0..aux_slots {
let coupon = cursor.read_u32_le().map_err(|_| {
Error::insufficient_data(format!(
- "expected {aux_count} aux coupons, failed at index
{i}",
+ "expected {aux_slots} HLL4 auxiliary slots, failed at
index {i}",
))
})?;
let coupon = Coupon(coupon);
+ if coupon.is_empty() && !compact {
+ continue;
+ }
let slot = coupon.slot() & ((1 << lg_config_k) - 1);
let value = coupon.value();
+ if coupon.is_empty() || aux.get(slot).is_some() {
+ return Err(Error::deserial(
+ "HLL4 auxiliary entries must be non-empty and unique",
+ ));
+ }
aux.insert(slot, value);
+ decoded_count += 1;
+ }
+ if decoded_count != aux_count as usize {
+ return Err(Error::deserial(format!(
+ "HLL4 auxiliary count is {aux_count}, decoded
{decoded_count}"
+ )));
}
aux_map = Some(aux);
}
diff --git a/datasketches/src/hll/array6.rs b/datasketches/src/hll/array6.rs
index 24ef2cb..69bfa22 100644
--- a/datasketches/src/hll/array6.rs
+++ b/datasketches/src/hll/array6.rs
@@ -178,10 +178,9 @@ impl Array6 {
/// Deserialize Array6 from HLL mode bytes
///
/// Expects full HLL preamble (40 bytes) followed by packed 6-bit data.
- pub fn deserialize(
+ pub fn deserialize_registers(
mut cursor: SketchSlice,
lg_config_k: u8,
- compact: bool,
ooo: bool,
) -> Result<Self, Error> {
let k = 1 << lg_config_k;
@@ -198,19 +197,26 @@ impl Array6 {
let num_zeros = cursor
.read_u32_le()
.map_err(insufficient_data("num_zeros"))?;
- let _aux_count = cursor
+ let aux_count = cursor
.read_u32_le()
- .map_err(insufficient_data("aux_count"))?; // always 0
+ .map_err(insufficient_data("aux_count"))?;
+ if num_zeros > k || aux_count != 0 {
+ return Err(Error::deserial(
+ "HLL6 zero count must not exceed k and auxiliary count must be
zero",
+ ));
+ }
+ if num_bytes > cursor.remaining().len() {
+ return Err(Error::insufficient_data(format!(
+ "HLL6 payload requires {num_bytes} bytes, got {}",
+ cursor.remaining().len()
+ )));
+ }
// Read packed byte array from offset HLL_BYTE_ARR_START
let mut data = vec![0u8; num_bytes];
- if !compact {
- cursor
- .read_exact(&mut data)
- .map_err(insufficient_data("data"))?;
- } else {
- cursor.advance(num_bytes as u64);
- }
+ cursor
+ .read_exact(&mut data)
+ .map_err(insufficient_data("data"))?;
// Create estimator and restore state
let mut estimator = HipEstimator::new(lg_config_k);
diff --git a/datasketches/src/hll/array8.rs b/datasketches/src/hll/array8.rs
index 45df645..0cc1e2a 100644
--- a/datasketches/src/hll/array8.rs
+++ b/datasketches/src/hll/array8.rs
@@ -252,10 +252,9 @@ impl Array8 {
/// Deserialize Array8 from HLL mode bytes
///
/// Expects full HLL preamble (40 bytes) followed by k bytes of data.
- pub fn deserialize(
+ pub fn deserialize_registers(
mut cursor: SketchSlice,
lg_config_k: u8,
- compact: bool,
ooo: bool,
) -> Result<Self, Error> {
let k = 1usize << lg_config_k;
@@ -271,19 +270,26 @@ impl Array8 {
let num_zeros = cursor
.read_u32_le()
.map_err(insufficient_data("num_zeros"))?;
- let _aux_count = cursor
+ let aux_count = cursor
.read_u32_le()
- .map_err(insufficient_data("aux_count"))?; // always 0
+ .map_err(insufficient_data("aux_count"))?;
+ if num_zeros as usize > k || aux_count != 0 {
+ return Err(Error::deserial(
+ "HLL8 zero count must not exceed k and auxiliary count must be
zero",
+ ));
+ }
+ if k > cursor.remaining().len() {
+ return Err(Error::insufficient_data(format!(
+ "HLL8 payload requires {k} bytes, got {}",
+ cursor.remaining().len()
+ )));
+ }
// Read byte array from offset HLL_BYTE_ARR_START
let mut data = vec![0u8; k];
- if !compact {
- cursor
- .read_exact(&mut data)
- .map_err(insufficient_data("data"))?;
- } else {
- cursor.advance(k as u64);
- }
+ cursor
+ .read_exact(&mut data)
+ .map_err(insufficient_data("data"))?;
// Create estimator and restore state
let mut estimator = HipEstimator::new(lg_config_k);
diff --git a/datasketches/src/hll/hash_set.rs b/datasketches/src/hll/hash_set.rs
index d2c6efb..66b4714 100644
--- a/datasketches/src/hll/hash_set.rs
+++ b/datasketches/src/hll/hash_set.rs
@@ -103,6 +103,20 @@ impl HashSet {
.read_u32_le()
.map_err(insufficient_data("coupon_count"))?;
let coupon_count = coupon_count as usize;
+ let array_size = 1usize << lg_arr;
+ if coupon_count >= array_size {
+ return Err(Error::deserial(format!(
+ "SET mode coupon count {coupon_count} must be below capacity
{array_size}"
+ )));
+ }
+ let read_count = if compact { coupon_count } else { array_size };
+ let required_bytes = read_count * size_of::<u32>();
+ if required_bytes > cursor.remaining().len() {
+ return Err(Error::insufficient_data(format!(
+ "SET mode coupons require {required_bytes} bytes, got {}",
+ cursor.remaining().len()
+ )));
+ }
if compact {
// Compact mode: only couponCount coupons are stored
@@ -116,11 +130,12 @@ impl HashSet {
})?;
hash_set.update(Coupon(coupon));
}
+ if hash_set.container.len() != coupon_count {
+ return Err(Error::deserial("SET mode contains duplicate
coupons"));
+ }
Ok(hash_set)
} else {
// Non-compact mode: full hash table with empty slots
- let array_size = 1 << lg_arr;
-
// Read entire hash table including empty slots
let mut coupons = vec![Coupon::EMPTY; array_size];
for (i, coupon) in coupons.iter_mut().enumerate() {
@@ -131,6 +146,11 @@ impl HashSet {
})?;
*coupon = Coupon(raw);
}
+ if coupons.iter().filter(|coupon| !coupon.is_empty()).count() !=
coupon_count {
+ return Err(Error::deserial(
+ "SET mode coupon count does not match occupied slots",
+ ));
+ }
Ok(Self {
container: Container::from_coupons(
diff --git a/datasketches/src/hll/list.rs b/datasketches/src/hll/list.rs
index 689e42c..afed99a 100644
--- a/datasketches/src/hll/list.rs
+++ b/datasketches/src/hll/list.rs
@@ -86,8 +86,25 @@ impl List {
// slots are available for future update() calls. In compact format
only
// coupon_count values are stored on disk, but memory must hold the
full capacity
// so the linear scan in update() can find an empty slot to insert
into.
- let array_size = 1 << lg_arr;
+ let array_size = 1usize << lg_arr;
+ if coupon_count > array_size {
+ return Err(Error::deserial(format!(
+ "LIST mode coupon count {coupon_count} exceeds capacity
{array_size}"
+ )));
+ }
+ if empty != (coupon_count == 0) {
+ return Err(Error::deserial(
+ "LIST mode empty flag and coupon count disagree",
+ ));
+ }
let read_count = if compact { coupon_count } else { array_size };
+ let required_bytes = read_count * size_of::<u32>();
+ if !empty && required_bytes > cursor.remaining().len() {
+ return Err(Error::insufficient_data(format!(
+ "LIST mode coupons require {required_bytes} bytes, got {}",
+ cursor.remaining().len()
+ )));
+ }
// Read coupons into the front of the full-sized array; remaining
slots stay Coupon::EMPTY.
let mut coupons = vec![Coupon::EMPTY; array_size];
diff --git a/datasketches/src/hll/serialization.rs
b/datasketches/src/hll/serialization.rs
index 30740a9..9f4eeeb 100644
--- a/datasketches/src/hll/serialization.rs
+++ b/datasketches/src/hll/serialization.rs
@@ -25,7 +25,9 @@ pub const SERIAL_VERSION: u8 = 1;
/// Flag indicating sketch is empty (no values inserted)
pub const EMPTY_FLAG_MASK: u8 = 4;
-/// Flag indicating compact serialization (no empty slots stored)
+/// Flag indicating compact coupon or HLL4 auxiliary storage.
+///
+/// HLL register arrays have the same layout in compact and updatable images.
pub const COMPACT_FLAG_MASK: u8 = 8;
/// Flag indicating out-of-order mode (HIP estimator invalid)
pub const OUT_OF_ORDER_FLAG_MASK: u8 = 16;
diff --git a/datasketches/src/hll/sketch.rs b/datasketches/src/hll/sketch.rs
index 21f5d53..befcc81 100644
--- a/datasketches/src/hll/sketch.rs
+++ b/datasketches/src/hll/sketch.rs
@@ -33,6 +33,7 @@ use crate::hll::HllType;
use crate::hll::RESIZE_DENOMINATOR;
use crate::hll::RESIZE_NUMERATOR;
use crate::hll::array4::Array4;
+use crate::hll::array4::AuxFormat;
use crate::hll::array6::Array6;
use crate::hll::array8::Array8;
use crate::hll::container::Container;
@@ -364,6 +365,11 @@ impl HllSketch {
)));
}
+ if lg_arr != 3 {
+ return Err(Error::deserial(format!(
+ "LIST mode lg_arr: expected 3, got {lg_arr}"
+ )));
+ }
let lg_arr = lg_arr as usize;
let coupon_count = state as usize;
let list = List::deserialize(cursor, lg_arr, coupon_count,
empty, compact)?;
@@ -377,6 +383,12 @@ impl HllSketch {
)));
}
+ let max_lg_arr = lg_config_k.saturating_sub(3);
+ if !(5..=max_lg_arr).contains(&lg_arr) {
+ return Err(Error::deserial(format!(
+ "SET mode lg_arr must be in [5, {max_lg_arr}], got
{lg_arr}"
+ )));
+ }
let lg_arr = lg_arr as usize;
let set = HashSet::deserialize(cursor, lg_arr, compact)?;
Mode::Set { set, hll_type }
@@ -391,13 +403,13 @@ impl HllSketch {
match hll_type {
HllType::Hll4 => {
- let cur_min = state;
- Array4::deserialize(cursor, cur_min, lg_config_k,
compact, ooo)
+ let aux = AuxFormat::from_header(compact, lg_arr);
+ Array4::deserialize(cursor, state, lg_config_k,
aux, ooo)
.map(Mode::Array4)?
}
- HllType::Hll6 => Array6::deserialize(cursor,
lg_config_k, compact, ooo)
+ HllType::Hll6 => Array6::deserialize_registers(cursor,
lg_config_k, ooo)
.map(Mode::Array6)?,
- HllType::Hll8 => Array8::deserialize(cursor,
lg_config_k, compact, ooo)
+ HllType::Hll8 => Array8::deserialize_registers(cursor,
lg_config_k, ooo)
.map(Mode::Array8)?,
}
}
diff --git a/datasketches/src/req/compactor.rs
b/datasketches/src/req/compactor.rs
index fc43a73..f614eec 100644
--- a/datasketches/src/req/compactor.rs
+++ b/datasketches/src/req/compactor.rs
@@ -28,20 +28,19 @@ use crate::req::nearest_even_section_size;
use crate::req::serialization::validate_compactor_state;
use crate::req::value::ReqValue;
-fn validate_deserialized_items<T: ReqValue>(items: &[T], sorted: bool) ->
Result<(), Error> {
- if items.iter().any(ReqValue::is_nan) {
- return Err(Error::deserial("REQ compactor contains a NaN item"));
- }
- if sorted
- && !items
- .windows(2)
- .all(|items| items[0].total_cmp(&items[1]).is_le())
- {
- return Err(Error::deserial(
- "REQ compactor claims to be sorted but its items are not",
- ));
+fn normalized_sort_state<T: ReqValue>(items: &[T], claimed_sorted: bool) ->
Result<bool, Error> {
+ let mut previous: Option<&T> = None;
+ let mut sorted = claimed_sorted;
+ for item in items {
+ if item.is_nan() {
+ return Err(Error::deserial("REQ compactor contains a NaN item"));
+ }
+ if sorted && previous.is_some_and(|previous|
previous.total_cmp(item).is_gt()) {
+ sorted = false;
+ }
+ previous = Some(item);
}
- Ok(())
+ Ok(sorted)
}
/// A compactor maintains items at a specific level of the REQ sketch.
@@ -285,7 +284,7 @@ where
self.items.truncate(self.items.len() - removed);
// Update state, then ensure enough sections (C++ order)
- self.state += 1;
+ self.state = self.state.wrapping_add(1);
self.ensure_enough_sections();
}
@@ -304,15 +303,6 @@ where
1u64 << self.lg_weight
}
- /// Returns the minimum stream length implied by this compactor's state.
- ///
- /// Each compaction removes at least two items of this level's weight.
Merging
- /// states with bitwise OR cannot make the result larger than their sum,
so the
- /// same lower bound holds for both updated and merged sketches.
- pub(super) fn minimum_stream_length(&self) -> Option<u64> {
- self.state.checked_mul(2)?.checked_mul(self.weight())
- }
-
// Private helper methods
fn ensure_enough_sections(&mut self) -> bool {
@@ -435,7 +425,7 @@ where
for _ in 0..num_items {
items.push(T::deserialize_value(cursor)?);
}
- validate_deserialized_items(&items, sorted)?;
+ let sorted = normalized_sort_state(&items, sorted)?;
Ok(Compactor::from_serialized_state(
lg_weight,
@@ -461,7 +451,7 @@ where
items: Vec<T>,
is_sorted: bool,
) -> Result<Self, Error> {
- validate_deserialized_items(&items, is_sorted)?;
+ let is_sorted = normalized_sort_state(&items, is_sorted)?;
let mut c = Self::new(0, k, rank_accuracy);
for item in items {
c.append(item);
diff --git a/datasketches/src/req/serialization.rs
b/datasketches/src/req/serialization.rs
index 0b5289b..2d56a5e 100644
--- a/datasketches/src/req/serialization.rs
+++ b/datasketches/src/req/serialization.rs
@@ -42,75 +42,21 @@ fn section_growth_threshold(num_sections: u8) ->
Option<u64> {
.and_then(|shift| 1u64.checked_shl(u32::from(shift)))
}
-fn reachable_section_doublings(state: u64, num_sections: u8) -> Option<u32> {
+fn has_reachable_section_count(state: u64, num_sections: u8) -> bool {
let mut sections = INITIAL_SECTIONS_PER_COMPACTOR;
- let mut doublings = 0;
while sections < num_sections {
- if state < section_growth_threshold(sections)? {
- return None;
+ let Some(threshold) = section_growth_threshold(sections) else {
+ return false;
+ };
+ if state < threshold {
+ return false;
}
- sections = sections.checked_mul(2)?;
- doublings += 1;
+ let Some(next) = sections.checked_mul(2) else {
+ return false;
+ };
+ sections = next;
}
- (sections == num_sections).then_some(doublings)
-}
-
-fn compatible_next_section_sizes(section_size_raw: f32) -> [Option<f32>; 2] {
- let single_precision = section_size_raw / std::f32::consts::SQRT_2;
- let widened = (f64::from(section_size_raw) / std::f64::consts::SQRT_2) as
f32;
- [
- (nearest_even_section_size(single_precision) >= u32::from(MIN_K))
- .then_some(single_precision),
- (nearest_even_section_size(section_size_raw) > u32::from(MIN_K)
- && nearest_even_section_size(widened) >= u32::from(MIN_K))
- .then_some(widened),
- ]
-}
-
-fn has_reachable_section_configuration(
- k: u16,
- state: u64,
- section_size_raw: f32,
- num_sections: u8,
-) -> bool {
- let Some(doublings) = reachable_section_doublings(state, num_sections)
else {
- return false;
- };
-
- // The wire format exposes the running floating-point section size. C++
- // updates it in single precision, while Java divides in double precision
- // before narrowing to f32. Accept either result at every step so a sketch
- // can be deserialized and continued in another implementation.
- let mut candidates = vec![k as f32];
- for _ in 0..doublings {
- let mut next = Vec::with_capacity(candidates.len() * 2);
- for candidate in candidates {
- for value in compatible_next_section_sizes(candidate)
- .into_iter()
- .flatten()
- {
- if !next.contains(&value) {
- next.push(value);
- }
- }
- }
- candidates = next;
- }
- if !candidates
- .iter()
- .any(|candidate| candidate.to_bits() == section_size_raw.to_bits())
- {
- return false;
- }
-
- // If every compatible producer would already have doubled the section
- // count at this state, the serialized configuration is stale.
- section_growth_threshold(num_sections).is_none_or(|threshold| {
- state < threshold
- || compatible_next_section_sizes(section_size_raw)
- .iter()
- .any(Option::is_none)
- })
+ sections == num_sections
}
pub(super) fn validate_compactor_state(
@@ -126,9 +72,13 @@ pub(super) fn validate_compactor_state(
"REQ compactor lg_weight {lg_weight} does not match level
{expected_lg_weight}"
)));
}
- if !has_reachable_section_configuration(k, state, section_size_raw,
num_sections) {
+ let section_size = nearest_even_section_size(section_size_raw);
+ if !section_size_raw.is_finite()
+ || !(u32::from(MIN_K)..=u32::from(k)).contains(§ion_size)
+ || !has_reachable_section_count(state, num_sections)
+ {
return Err(Error::deserial(format!(
- "REQ compactor section configuration is not reachable (k={k},
state={state}, section_size_raw={section_size_raw},
num_sections={num_sections})"
+ "REQ compactor layout is invalid (k={k}, state={state},
section_size_raw={section_size_raw}, num_sections={num_sections})"
)));
}
Ok(())
diff --git a/datasketches/src/req/sketch.rs b/datasketches/src/req/sketch.rs
index bec4d8f..d3a6cc3 100644
--- a/datasketches/src/req/sketch.rs
+++ b/datasketches/src/req/sketch.rs
@@ -718,14 +718,6 @@ impl<T: ReqValue> ReqSketch<T> {
));
}
- if compactors.iter().any(|compactor| {
- compactor
- .minimum_stream_length()
- .is_none_or(|minimum_n| minimum_n > n)
- }) {
- return Err(Error::deserial("REQ compactor state exceeds stream
length"));
- }
-
let (retained_count, nominal_capacity, weighted_count) = compactors
.iter()
.try_fold(
diff --git a/datasketches/src/thetafamily/theta/sketch.rs
b/datasketches/src/thetafamily/theta/sketch.rs
index d1870cb..4293e7c 100644
--- a/datasketches/src/thetafamily/theta/sketch.rs
+++ b/datasketches/src/thetafamily/theta/sketch.rs
@@ -724,6 +724,15 @@ impl CompactThetaSketch {
num_entries: usize,
theta: u64,
) -> Result<Vec<u64>, Error> {
+ let required_bytes = num_entries
+ .checked_mul(size_of::<u64>())
+ .ok_or_else(|| Error::deserial("Theta entry payload length
overflows"))?;
+ if required_bytes > cursor.remaining().len() {
+ return Err(Error::insufficient_data(format!(
+ "Theta entries require {required_bytes} bytes, got {}",
+ cursor.remaining().len()
+ )));
+ }
let mut entries = Vec::with_capacity(num_entries);
for _ in 0..num_entries {
let hash =
cursor.read_u64_le().map_err(insufficient_data("entries"))?;
@@ -900,6 +909,11 @@ impl CompactThetaSketch {
) -> Result<Self, Error> {
let entry_bits =
cursor.read_u8().map_err(insufficient_data("entry_bits"))?;
let num_entries_bytes =
cursor.read_u8().map_err(insufficient_data("num_entries"))?;
+ if num_entries_bytes > size_of::<u32>() as u8 {
+ return Err(Error::deserial(format!(
+ "Theta entry count uses too many bytes: {num_entries_bytes}"
+ )));
+ }
let flags = cursor.read_u8().map_err(insufficient_data("flags"))?;
let seed_hash = cursor
.read_u16_le()
@@ -929,6 +943,22 @@ impl CompactThetaSketch {
.map_err(insufficient_data("num_entries_byte"))?;
num_entries |= (entry_count_byte as usize) << ((i as usize) << 3);
}
+ if num_entries > 0 && !(1..=63).contains(&entry_bits) {
+ return Err(Error::deserial(format!(
+ "Theta entry width must be in [1, 63], got {entry_bits}"
+ )));
+ }
+ let required_bytes = num_entries
+ .checked_mul(entry_bits as usize)
+ .and_then(|bits| bits.checked_add(7))
+ .map(|bits| bits / 8)
+ .ok_or_else(|| Error::deserial("Theta compressed payload length
overflows"))?;
+ if required_bytes > cursor.remaining().len() {
+ return Err(Error::insufficient_data(format!(
+ "Theta compressed entries require {required_bytes} bytes, got
{}",
+ cursor.remaining().len()
+ )));
+ }
// unpack blocks of BLOCK_WIDTH deltas
let mut i = 0usize;
@@ -961,7 +991,9 @@ impl CompactThetaSketch {
// undo deltas
let mut previous = 0;
for e in &mut entries {
- *e += previous;
+ *e = e
+ .checked_add(previous)
+ .ok_or_else(|| Error::deserial("Theta entry delta
overflows"))?;
previous = *e;
if *e == 0 || *e >= theta {
return Err(Error::deserial("corrupted: invalid retained hash
value"));
diff --git a/datasketches/src/thetafamily/tuple/sketch.rs
b/datasketches/src/thetafamily/tuple/sketch.rs
index 8cb9c53..fd1c8fb 100644
--- a/datasketches/src/thetafamily/tuple/sketch.rs
+++ b/datasketches/src/thetafamily/tuple/sketch.rs
@@ -689,6 +689,15 @@ impl<S> CompactTupleSketch<S> {
n
};
+ let required_hash_bytes = num_entries
+ .checked_mul(size_of::<u64>())
+ .ok_or_else(|| Error::deserial("Tuple entry payload length
overflows"))?;
+ if required_hash_bytes > cursor.remaining().len() {
+ return Err(Error::insufficient_data(format!(
+ "Tuple entry hashes require at least {required_hash_bytes}
bytes, got {}",
+ cursor.remaining().len()
+ )));
+ }
let mut entries = Vec::with_capacity(num_entries);
for _ in 0..num_entries {
let hash = cursor
diff --git a/datasketches/tests/cpc_test/deserialize.rs
b/datasketches/tests/cpc_test/deserialize.rs
new file mode 100644
index 0000000..3a77577
--- /dev/null
+++ b/datasketches/tests/cpc_test/deserialize.rs
@@ -0,0 +1,115 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements. See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership. The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License. You may obtain a copy of the License at
+//
+// http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied. See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+//! Regression tests for deserializing malformed CPC sketches.
+
+use datasketches::cpc::CpcSketch;
+
+/// Builds a valid serialized sketch that exercises a particular CPC flavor.
+fn valid_bytes(lg_k: u8, n: u64) -> Vec<u8> {
+ let mut sketch = CpcSketch::new(lg_k);
+ for i in 0..n {
+ sketch.update(i);
+ }
+ sketch.serialize()
+}
+
+/// Each non-empty entry lands in a different CPC flavor.
+const CASES: &[(u8, u64)] = &[
+ (4, 0),
+ (4, 3),
+ (6, 20),
+ (8, 200),
+ (8, 2000),
+ (10, 800),
+ (10, 8000),
+];
+
+#[test]
+fn truncated_compressed_streams_return_errors() {
+ let bytes = valid_bytes(10, 8_000);
+ for len in 0..bytes.len() {
+ assert!(
+ CpcSketch::deserialize(&bytes[..len]).is_err(),
+ "accepted truncation at {len} bytes"
+ );
+ }
+}
+
+#[test]
+fn targeted_corruptions_return_err() {
+ // A sliding-flavor sketch drives the pair/window decoders and the
offset/permutation logic.
+ let base = valid_bytes(10, 8000);
+
+ // Layout: [preamble_ints, serial_version, family, lg_k,
first_interesting_column, flags,
+ // seed_hash(2), num_coupons(4), ...]. Corrupting the num_coupons
field makes the
+ // decoders read past the compressed buffer / compute an
out-of-range window offset.
+ let mut num_coupons_hi = base.clone();
+ num_coupons_hi[11] = 0xff; // enormous num_coupons
+ assert!(CpcSketch::deserialize(&num_coupons_hi).is_err());
+
+ // Flipping the flags byte makes the declared flavor inconsistent with the
stored data.
+ let mut bad_flags = base.clone();
+ bad_flags[5] ^= 0xff;
+ assert!(CpcSketch::deserialize(&bad_flags).is_err());
+
+ // These payload edits previously reached panicking decoder and pair-table
paths.
+ let mut bad_payload = base.clone();
+ let last = bad_payload.len() - 1;
+ bad_payload[last] = bad_payload[last].wrapping_add(1);
+ bad_payload[last - 3] ^= 0xa5;
+ let _ = CpcSketch::deserialize(&bad_payload);
+
+ // A sparse sketch whose declared entry count exceeds its data words must
be rejected up front.
+ let sparse = valid_bytes(8, 200);
+ let mut inflated = sparse.clone();
+ // num_coupons is a u32 at offset 8; inflate it far beyond the coupon
space.
+ inflated[10] = 0xff;
+ inflated[11] = 0xff;
+ assert!(CpcSketch::deserialize(&inflated).is_err());
+}
+
+#[test]
+fn valid_sketches_round_trip_unchanged() {
+ for &(lg_k, n) in CASES {
+ let mut sketch = CpcSketch::new(lg_k);
+ for i in 0..n {
+ sketch.update(i);
+ }
+ let bytes = sketch.serialize();
+
+ let restored = CpcSketch::deserialize(&bytes).unwrap_or_else(|e| {
+ panic!("valid sketch (lg_k={lg_k}, n={n}) failed to deserialize:
{e}")
+ });
+
+ assert_eq!(
+ sketch.estimate(),
+ restored.estimate(),
+ "estimate changed after round-trip (lg_k={lg_k}, n={n})"
+ );
+ assert_eq!(
+ sketch.num_coupons(),
+ restored.num_coupons(),
+ "num_coupons changed after round-trip (lg_k={lg_k}, n={n})"
+ );
+ assert_eq!(
+ bytes,
+ restored.serialize(),
+ "re-serialized bytes changed after round-trip (lg_k={lg_k}, n={n})"
+ );
+ }
+}
diff --git a/datasketches/tests/cpc_test/main.rs
b/datasketches/tests/cpc_test/main.rs
index 0583ff1..7b98ba9 100644
--- a/datasketches/tests/cpc_test/main.rs
+++ b/datasketches/tests/cpc_test/main.rs
@@ -15,6 +15,7 @@
// specific language governing permissions and limitations
// under the License.
+mod deserialize;
mod union;
mod update;
mod wrapper;
diff --git a/datasketches/tests/serde_tests/bloom.rs
b/datasketches/tests/serde_tests/bloom.rs
index b8f6b38..708c7f7 100644
--- a/datasketches/tests/serde_tests/bloom.rs
+++ b/datasketches/tests/serde_tests/bloom.rs
@@ -20,7 +20,6 @@ use std::path::PathBuf;
use datasketches::bloom::BloomFilter;
use datasketches::bloom::BloomFilterBuilder;
-use datasketches::error::ErrorKind;
use googletest::assert_that;
use googletest::prelude::gt;
@@ -178,7 +177,7 @@ fn test_go_compatibility() {
}
#[test]
-fn test_inconsistent_num_bits_set_is_rejected() {
+fn test_cached_num_bits_set_is_validated_or_recomputed() {
const NUM_BITS_SET_OFFSET: usize = 24;
let mut filter = BloomFilterBuilder::with_accuracy(100, 0.01).build();
@@ -192,24 +191,27 @@ fn test_inconsistent_num_bits_set_is_rejected() {
bytes[NUM_BITS_SET_OFFSET..NUM_BITS_SET_OFFSET + size_of::<u64>()]
.copy_from_slice(&serialized_count.to_le_bytes());
- let err = BloomFilter::deserialize(&bytes).unwrap_err();
- assert_eq!(err.kind(), ErrorKind::InvalidData);
+ assert!(BloomFilter::deserialize(&bytes).is_err());
}
+
+ let mut dirty = filter.serialize();
+ dirty[NUM_BITS_SET_OFFSET..NUM_BITS_SET_OFFSET + size_of::<u64>()]
+ .copy_from_slice(&u64::MAX.to_le_bytes());
+ let restored = BloomFilter::deserialize(&dirty).unwrap();
+ assert_eq!(restored.bits_used(), actual_bits_set);
+ assert!(restored.contains(&"apple"));
+ assert!(restored.contains(&"banana"));
}
#[test]
-fn test_dirty_num_bits_set_is_recomputed() {
- const NUM_BITS_SET_OFFSET: usize = 24;
+fn test_nonempty_payload_length_is_checked_before_allocating() {
+ const NUM_LONGS_OFFSET: usize = 16;
let mut filter = BloomFilterBuilder::with_accuracy(100, 0.01).build();
filter.insert("apple");
- filter.insert("banana");
let mut bytes = filter.serialize();
- bytes[NUM_BITS_SET_OFFSET..NUM_BITS_SET_OFFSET + size_of::<u64>()]
- .copy_from_slice(&u64::MAX.to_le_bytes());
+ bytes[NUM_LONGS_OFFSET..NUM_LONGS_OFFSET + size_of::<i32>()]
+ .copy_from_slice(&i32::MAX.to_le_bytes());
- let restored = BloomFilter::deserialize(&bytes).unwrap();
- assert_eq!(restored.bits_used(), filter.bits_used());
- assert!(restored.contains(&"apple"));
- assert!(restored.contains(&"banana"));
+ assert!(BloomFilter::deserialize(&bytes).is_err());
}
diff --git a/datasketches/tests/serde_tests/hll.rs
b/datasketches/tests/serde_tests/hll.rs
index f491489..b3f5207 100644
--- a/datasketches/tests/serde_tests/hll.rs
+++ b/datasketches/tests/serde_tests/hll.rs
@@ -125,6 +125,81 @@ fn test_update_after_deserialize_list_mode() {
}
}
+#[test]
+fn coupon_mode_sizes_are_validated_before_allocating() {
+ let mut list = HllSketch::new(12, HllType::Hll8);
+ list.update(1_u64);
+ let mut invalid_list_size = list.serialize();
+ invalid_list_size[4] = u8::MAX;
+ assert!(HllSketch::deserialize(&invalid_list_size).is_err());
+
+ let mut invalid_list_count = list.serialize();
+ invalid_list_count[6] = u8::MAX;
+ assert!(HllSketch::deserialize(&invalid_list_count).is_err());
+
+ let mut set = HllSketch::new(12, HllType::Hll8);
+ for value in 0..10 {
+ set.update(value);
+ }
+ let mut invalid_set_size = set.serialize();
+ invalid_set_size[4] = u8::MAX;
+ assert!(HllSketch::deserialize(&invalid_set_size).is_err());
+
+ let mut invalid_set_count = set.serialize();
+ invalid_set_count[8..12].copy_from_slice(&u32::MAX.to_le_bytes());
+ assert!(HllSketch::deserialize(&invalid_set_count).is_err());
+}
+
+#[test]
+fn hll_mode_round_trip_preserves_registers_and_rejects_truncation() {
+ for hll_type in [HllType::Hll4, HllType::Hll6, HllType::Hll8] {
+ let mut sketch = HllSketch::new(12, hll_type);
+ for value in 0..10_000 {
+ sketch.update(value);
+ }
+
+ let bytes = sketch.serialize();
+ let restored = HllSketch::deserialize(&bytes).unwrap();
+ assert_eq!(restored, sketch, "hll_type: {hll_type:?}");
+ assert!(HllSketch::deserialize(&bytes[..bytes.len() - 1]).is_err());
+ }
+}
+
+#[test]
+fn hll4_updatable_aux_table_matches_compact_image() {
+ let path = serialization_test_data("java_generated_files",
"hll4_n100000_java.sk");
+ let compact = fs::read(path).unwrap();
+ let lg_k = compact[3];
+ let lg_arr = compact[4];
+ let aux_count = u32::from_le_bytes(compact[36..40].try_into().unwrap()) as
usize;
+ let aux_start = 40 + (1usize << lg_k) / 2;
+ assert!(aux_count > 0);
+
+ let mut aux_table = vec![0_u32; 1usize << lg_arr];
+ let table_mask = aux_table.len() as u32 - 1;
+ for bytes in compact[aux_start..aux_start + aux_count *
size_of::<u32>()].chunks_exact(4) {
+ let coupon = u32::from_le_bytes(bytes.try_into().unwrap());
+ let slot = coupon & ((1_u32 << 26) - 1);
+ let stride = (slot >> lg_arr) | 1;
+ let mut probe = slot & table_mask;
+ while aux_table[probe as usize] != 0 {
+ probe = (probe + stride) & table_mask;
+ }
+ aux_table[probe as usize] = coupon;
+ }
+
+ let mut updatable = compact[..aux_start].to_vec();
+ updatable[5] &= !8; // clear the compact flag
+ for coupon in aux_table {
+ updatable.extend_from_slice(&coupon.to_le_bytes());
+ }
+
+ assert_eq!(
+ HllSketch::deserialize(&updatable).unwrap(),
+ HllSketch::deserialize(&compact).unwrap()
+ );
+}
+
#[test]
fn test_serialized_bytes_match_reference_files_for_coupon_modes() {
fn serialized_mode_name(bytes: &[u8]) -> &'static str {
diff --git a/datasketches/tests/serde_tests/req.rs
b/datasketches/tests/serde_tests/req.rs
index 445fbf3..d9d8acf 100644
--- a/datasketches/tests/serde_tests/req.rs
+++ b/datasketches/tests/serde_tests/req.rs
@@ -357,7 +357,7 @@ fn deserialize_accepts_java_minimum_section_schedule() {
}
#[test]
-fn deserialize_rejects_capacity_changing_float_drift() {
+fn deserialize_accepts_safe_section_size_drift() {
let mut bytes = estimation_image(10, 2_562);
let raw_offset = ESTIMATION_COMPACTOR_OFFSET + SECTION_SIZE_RAW_OFFSET;
let raw_bits = u32::from_le_bytes(bytes[raw_offset..raw_offset +
4].try_into().unwrap());
@@ -367,11 +367,16 @@ fn deserialize_rejects_capacity_changing_float_drift() {
// One ULP below 5.0 rounds to a section size of 4 rather than 6.
bytes[raw_offset..raw_offset + 4].copy_from_slice(&(raw_bits -
1).to_le_bytes());
- assert_invalid_data(&bytes);
+ let mut sketch = ReqSketch::<f32>::deserialize(&bytes).unwrap();
+ sketch.update(2_563.0);
+ assert_that!(
+ ReqSketch::<f32>::deserialize(&sketch.serialize()),
+ ok(anything())
+ );
}
#[test]
-fn deserialize_rejects_state_inconsistent_with_stream_length() {
+fn deserialize_accepts_opaque_compactor_state() {
let mut bytes = estimation_image(12, 1_000);
let compactor = ESTIMATION_COMPACTOR_OFFSET;
let state = 501u64;
@@ -384,11 +389,16 @@ fn
deserialize_rejects_state_inconsistent_with_stream_length() {
bytes[compactor + SECTION_SIZE_RAW_OFFSET..compactor +
SECTION_SIZE_RAW_OFFSET + 4]
.copy_from_slice(&raw.to_le_bytes());
bytes[compactor + NUM_SECTIONS_OFFSET] = 12;
- assert_invalid_data(&bytes);
+ let mut sketch = ReqSketch::<f32>::deserialize(&bytes).unwrap();
+ sketch.update(1_001.0);
+ assert_that!(
+ ReqSketch::<f32>::deserialize(&sketch.serialize()),
+ ok(anything())
+ );
}
#[test]
-fn deserialize_rejects_complementary_states_that_overflow_on_merge() {
+fn complementary_compactor_states_do_not_overflow_on_merge() {
let items: Vec<f32> = (0..192).map(|item| item as f32).collect();
let mut raw = 12.0f32;
for _ in 0..4 {
@@ -397,7 +407,7 @@ fn
deserialize_rejects_complementary_states_that_overflow_on_merge() {
let states = [0xAAAA_AAAA_AAAA_AAAAu64, 0x5555_5555_5555_5555u64];
assert_eq!(states[0] | states[1], u64::MAX);
- for state in states {
+ let mut sketches = states.map(|state| {
let mut bytes = exact_image(12, &items);
let compactor = EXACT_COMPACTOR_OFFSET;
bytes[compactor + STATE_OFFSET..compactor + STATE_OFFSET + 8]
@@ -405,15 +415,24 @@ fn
deserialize_rejects_complementary_states_that_overflow_on_merge() {
bytes[compactor + SECTION_SIZE_RAW_OFFSET..compactor +
SECTION_SIZE_RAW_OFFSET + 4]
.copy_from_slice(&raw.to_le_bytes());
bytes[compactor + NUM_SECTIONS_OFFSET] = 48;
- assert_invalid_data(&bytes);
- }
+ ReqSketch::<f32>::deserialize(&bytes).unwrap()
+ });
+ let (left, right) = sketches.split_at_mut(1);
+ left[0].merge(&right[0]).unwrap();
}
#[test]
-fn deserialize_rejects_false_sorted_claim_and_nan() {
+fn deserialize_normalizes_false_sorted_claim_and_rejects_nan() {
let mut unsorted = exact_image(12, &[3.0, 4.0, 5.0, 1.0, 2.0]);
unsorted[3] |= FLAG_LEVEL_ZERO_SORTED;
- assert_invalid_data(&unsorted);
+ let sketch = ReqSketch::<f32>::deserialize(&unsorted).unwrap();
+ assert_eq!(
+ sketch.rank(&2.0, SearchCriteria::Inclusive).unwrap(),
+ sketch
+ .sorted_view()
+ .rank(&2.0, SearchCriteria::Inclusive)
+ .unwrap(),
+ );
let nan = exact_image(12, &[1.0, 2.0, f32::NAN, 4.0, 5.0]);
assert_invalid_data(&nan);
diff --git a/datasketches/tests/serde_tests/theta.rs
b/datasketches/tests/serde_tests/theta.rs
index b09ec8f..e1ee343 100644
--- a/datasketches/tests/serde_tests/theta.rs
+++ b/datasketches/tests/serde_tests/theta.rs
@@ -184,6 +184,29 @@ fn malformed_input_is_rejected() {
assert_eq!(err.kind(), ErrorKind::InvalidData);
}
+#[test]
+fn declared_entry_payload_is_checked_before_allocating() {
+ let mut uncompressed = serialize_v2_exact(&[1]);
+ uncompressed[8..12].copy_from_slice(&u32::MAX.to_le_bytes());
+ assert!(CompactThetaSketch::deserialize(&uncompressed).is_err());
+
+ let mut sketch = ThetaSketchBuilder::default().lg_k(5).build();
+ for value in 0..5000 {
+ sketch.update(value);
+ }
+ let compressed = sketch.compact(true).serialize_compressed();
+
+ let mut invalid_entry_width = compressed.clone();
+ invalid_entry_width[3] = 64;
+ assert!(CompactThetaSketch::deserialize(&invalid_entry_width).is_err());
+
+ let mut oversized_entry_count = compressed;
+ let count_offset = usize::from(oversized_entry_count[0]) *
size_of::<u64>();
+ oversized_entry_count[4] = 4;
+ oversized_entry_count[count_offset..count_offset +
size_of::<u32>()].fill(u8::MAX);
+ assert!(CompactThetaSketch::deserialize(&oversized_entry_count).is_err());
+}
+
#[test]
fn test_v2_exact_non_empty_compatibility() {
let entries = [1, 7, 42];
diff --git a/datasketches/tests/serde_tests/tuple.rs
b/datasketches/tests/serde_tests/tuple.rs
index b5c36f9..33103f0 100644
--- a/datasketches/tests/serde_tests/tuple.rs
+++ b/datasketches/tests/serde_tests/tuple.rs
@@ -143,3 +143,16 @@ fn malformed_input_is_rejected() {
let err =
CompactTupleSketch::<u64>::deserialize(&wrong_family).unwrap_err();
assert_eq!(err.kind(), ErrorKind::InvalidData);
}
+
+#[test]
+fn declared_entry_payload_is_checked_before_allocating() {
+ let mut sketch =
TupleSketchBuilder::new(DefaultUpdatePolicy::<u64>::default()).build();
+ for value in 0..100 {
+ sketch.update(value, 1);
+ }
+ let mut bytes = sketch.compact(true).serialize();
+ assert!(bytes[0] > 1);
+ bytes[8..12].copy_from_slice(&u32::MAX.to_le_bytes());
+
+ assert!(CompactTupleSketch::<u64>::deserialize(&bytes).is_err());
+}
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]