diff --git a/datasketches/src/bloom/sketch.rs b/datasketches/src/bloom/sketch.rs index 081f6d4..64af2e5 100644 --- a/datasketches/src/bloom/sketch.rs +++ b/datasketches/src/bloom/sketch.rs @@ -471,23 +471,22 @@ impl BloomFilter { .read_u64_le() .map_err(insufficient_data("num_bits_set"))?; + let mut counted_bits_set = 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; } - // Handle "dirty" state: 0xFFFFFFFFFFFFFFFF indicates bits need recounting + // Handle "dirty" state: 0xFFFFFFFFFFFFFFFF indicates bits need recounting. const DIRTY_BITS_VALUE: u64 = 0xFFFFFFFFFFFFFFFF; if raw_num_bits_set == DIRTY_BITS_VALUE { - num_bits_set = bit_array.iter().map(|w| w.count_ones() as u64).sum(); + num_bits_set = counted_bits_set; } else { - let raw_num_words_set = raw_num_bits_set.div_ceil(64) as usize; - if raw_num_words_set > num_words { + if raw_num_bits_set != counted_bits_set { return Err(Error::deserial(format!( - "invalid num_bits_set: expected <= {}, got {}", - num_words * 64, - raw_num_bits_set + "invalid num_bits_set: expected {counted_bits_set}, got {raw_num_bits_set}", ))); } num_bits_set = raw_num_bits_set; diff --git a/datasketches/tests/serde_tests/bloom.rs b/datasketches/tests/serde_tests/bloom.rs index 6e406ec..ac5df83 100644 --- a/datasketches/tests/serde_tests/bloom.rs +++ b/datasketches/tests/serde_tests/bloom.rs @@ -19,6 +19,8 @@ use std::fs; use std::path::PathBuf; use datasketches::bloom::BloomFilter; +use datasketches::bloom::BloomFilterBuilder; +use datasketches::error::ErrorKind; use crate::serialization_test_data; @@ -175,3 +177,40 @@ fn test_go_compatibility() { test_bloom_filter_file(path, n, num_hashes); } } + +#[test] +fn test_inconsistent_num_bits_set_is_rejected() { + const NUM_BITS_SET_OFFSET: usize = 24; + + let mut filter = BloomFilterBuilder::with_accuracy(100, 0.01).build(); + filter.insert("apple"); + filter.insert("banana"); + let actual_bits_set = filter.bits_used(); + assert!(actual_bits_set > 1); + + for serialized_count in [0, actual_bits_set - 1, actual_bits_set + 1] { + let mut bytes = filter.serialize(); + bytes[NUM_BITS_SET_OFFSET..NUM_BITS_SET_OFFSET + size_of::()] + .copy_from_slice(&serialized_count.to_le_bytes()); + + let err = BloomFilter::deserialize(&bytes).unwrap_err(); + assert_eq!(err.kind(), ErrorKind::InvalidData); + } +} + +#[test] +fn test_dirty_num_bits_set_is_recomputed() { + const NUM_BITS_SET_OFFSET: usize = 24; + + 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::()] + .copy_from_slice(&u64::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")); +}