use std::hash::Hash;
use std::hash::Hasher;
use crate::codec::SketchBytes;
use crate::codec::SketchSlice;
use crate::codec::assert::ensure_preamble_longs_in_range;
use crate::codec::assert::ensure_serial_version_is;
use crate::codec::assert::insufficient_data;
use crate::codec::family::Family;
use crate::error::Error;
use crate::hash::DEFAULT_UPDATE_SEED;
use crate::hash::XxHash64;
const SERIAL_VERSION: u8 = 1;
const EMPTY_FLAG_MASK: u8 = 1 << 2;
#[derive(Debug, Clone, PartialEq)]
pub struct BloomFilter {
seed: u64,
num_hashes: u16,
num_bits_set: u64,
bit_array: Box<[u64]>,
}
impl BloomFilter {
pub fn contains<T: Hash>(&self, item: &T) -> bool {
if self.is_empty() {
return false;
}
let (h0, h1) = self.compute_hash(item);
self.check_bits(h0, h1)
}
pub fn contains_and_insert<T: Hash>(&mut self, item: &T) -> bool {
let (h0, h1) = self.compute_hash(item);
let was_present = self.check_bits(h0, h1);
self.set_bits(h0, h1);
was_present
}
pub fn insert<T: Hash>(&mut self, item: T) {
let (h0, h1) = self.compute_hash(&item);
self.set_bits(h0, h1);
}
pub fn reset(&mut self) {
self.bit_array.fill(0);
self.num_bits_set = 0
}
pub fn union(&mut self, other: &BloomFilter) -> Result<(), Error> {
if !self.is_compatible(other) {
return Err(Error::invalid_argument(
"Bloom filters must have matching capacity, number of hashes, and seed",
));
}
let mut num_bits_set = 0;
for (word, other_word) in self.bit_array.iter_mut().zip(&other.bit_array) {
*word |= *other_word;
num_bits_set += word.count_ones() as u64;
}
self.num_bits_set = num_bits_set;
Ok(())
}
pub fn intersect(&mut self, other: &BloomFilter) -> Result<(), Error> {
if !self.is_compatible(other) {
return Err(Error::invalid_argument(
"Bloom filters must have matching capacity, number of hashes, and seed",
));
}
let mut num_bits_set = 0;
for (word, other_word) in self.bit_array.iter_mut().zip(&other.bit_array) {
*word &= *other_word;
num_bits_set += word.count_ones() as u64;
}
self.num_bits_set = num_bits_set;
Ok(())
}
pub fn invert(&mut self) {
for word in &mut self.bit_array {
*word = !*word;
}
self.num_bits_set = self.capacity() as u64 - self.num_bits_set;
}
pub fn is_empty(&self) -> bool {
self.num_bits_set == 0
}
pub fn bits_used(&self) -> u64 {
self.num_bits_set
}
pub fn capacity(&self) -> usize {
self.bit_array.len() * 64
}
pub fn num_hashes(&self) -> u16 {
self.num_hashes
}
pub fn seed(&self) -> u64 {
self.seed
}
pub fn load_factor(&self) -> f64 {
self.num_bits_set as f64 / self.capacity() as f64
}
pub fn estimated_fpp(&self) -> f64 {
let k = self.num_hashes as f64;
let load = self.load_factor();
load.powf(k)
}
pub fn is_compatible(&self, other: &Self) -> bool {
self.bit_array.len() == other.bit_array.len()
&& self.num_hashes == other.num_hashes
&& self.seed == other.seed
}
pub fn serialize(&self) -> Vec<u8> {
let is_empty = self.is_empty();
let preamble_longs = if is_empty {
Family::BLOOMFILTER.min_pre_longs
} else {
Family::BLOOMFILTER.max_pre_longs
};
let capacity = 8 * preamble_longs as usize
+ if is_empty {
0
} else {
self.bit_array.len() * 8
};
let mut bytes = SketchBytes::with_capacity(capacity);
bytes.write_u8(preamble_longs); bytes.write_u8(SERIAL_VERSION); bytes.write_u8(Family::BLOOMFILTER.id); bytes.write_u8(if is_empty { EMPTY_FLAG_MASK } else { 0 }); bytes.write_u16_le(self.num_hashes); bytes.write_u16_le(0);
bytes.write_u64_le(self.seed);
let num_longs = self.bit_array.len() as i32;
bytes.write_i32_le(num_longs);
bytes.write_u32_le(0);
if !is_empty {
bytes.write_u64_le(self.num_bits_set);
for &word in &self.bit_array {
bytes.write_u64_le(word);
}
}
bytes.into_bytes()
}
pub fn deserialize(bytes: &[u8]) -> Result<Self, Error> {
let mut cursor = SketchSlice::new(bytes);
let preamble_longs = cursor
.read_u8()
.map_err(insufficient_data("preamble_longs"))?;
let serial_version = cursor
.read_u8()
.map_err(insufficient_data("serial_version"))?;
let family_id = cursor.read_u8().map_err(insufficient_data("family_id"))?;
let flags = cursor.read_u8().map_err(insufficient_data("flags"))?;
Family::BLOOMFILTER.validate_id(family_id)?;
ensure_serial_version_is(SERIAL_VERSION, serial_version)?;
ensure_preamble_longs_in_range(
Family::BLOOMFILTER.min_pre_longs..=Family::BLOOMFILTER.max_pre_longs,
preamble_longs,
)?;
let is_empty = (flags & EMPTY_FLAG_MASK) != 0;
let num_hashes = cursor
.read_u16_le()
.map_err(insufficient_data("num_hashes"))?;
if num_hashes == 0 || num_hashes > i16::MAX as u16 {
return Err(Error::deserial(format!(
"invalid num_hashes: expected [1, {}], got {}",
i16::MAX,
num_hashes
)));
}
let _unused = cursor
.read_u16_le()
.map_err(insufficient_data("unused_header"))?;
let seed = cursor.read_u64_le().map_err(insufficient_data("seed"))?;
let num_longs = cursor
.read_i32_le()
.map_err(insufficient_data("num_longs"))?;
let _unused = cursor.read_u32_le().map_err(insufficient_data("unused"))?;
if num_longs <= 0 {
return Err(Error::deserial(format!(
"invalid num_longs: expected at least 1, got {}",
num_longs
)));
}
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 {
0
} else {
let serialized_num_bits_set = cursor
.read_u64_le()
.map_err(insufficient_data("num_bits_set"))?;
let mut count = 0;
for word in &mut bit_array {
*word = cursor
.read_u64_le()
.map_err(insufficient_data("bit_array"))?;
count += word.count_ones() as u64;
}
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,
num_hashes,
num_bits_set,
bit_array,
})
}
fn compute_hash<T: Hash>(&self, item: &T) -> (u64, u64) {
let mut hasher = XxHash64::with_seed(self.seed);
item.hash(&mut hasher);
let h0 = hasher.finish();
let mut hasher = XxHash64::with_seed(h0);
item.hash(&mut hasher);
let h1 = hasher.finish();
(h0, h1)
}
fn check_bits(&self, h0: u64, h1: u64) -> bool {
for i in 1..=self.num_hashes {
let bit_index = self.compute_bit_index(h0, h1, i);
if !self.get_bit(bit_index) {
return false;
}
}
true
}
fn set_bits(&mut self, h0: u64, h1: u64) {
for i in 1..=self.num_hashes {
let bit_index = self.compute_bit_index(h0, h1, i);
self.set_bit(bit_index);
}
}
fn compute_bit_index(&self, h0: u64, h1: u64, i: u16) -> usize {
let hash = h0.wrapping_add(u64::from(i).wrapping_mul(h1)) as usize;
(hash >> 1) % self.capacity()
}
fn get_bit(&self, bit_index: usize) -> bool {
let word_index = bit_index >> 6; let bit_offset = bit_index & 63; let mask = 1u64 << bit_offset;
(self.bit_array[word_index] & mask) != 0
}
fn set_bit(&mut self, bit_index: usize) {
let word_index = bit_index >> 6; let bit_offset = bit_index & 63; let mask = 1u64 << bit_offset;
if (self.bit_array[word_index] & mask) == 0 {
self.bit_array[word_index] |= mask;
self.num_bits_set += 1;
}
}
pub fn estimated_size(&self) -> usize {
size_of::<Self>() + self.bit_array.len() * size_of::<u64>()
}
}
#[derive(Debug, Clone)]
pub struct BloomFilterBuilder {
mode: BloomFilterBuilderMode,
seed: u64,
}
#[derive(Debug, Clone)]
enum BloomFilterBuilderMode {
Accuracy { max_items: u64, fpp: f64 },
Size { num_bits: u64, num_hashes: u16 },
}
impl BloomFilterBuilder {
const MIN_NUM_BITS: u64 = 1;
const MAX_NUM_BITS: u64 = (i32::MAX as u64 - Family::BLOOMFILTER.max_pre_longs as u64) * 64;
const MIN_NUM_HASHES: u16 = 1;
const MAX_NUM_HASHES: u16 = i16::MAX as u16;
pub fn with_accuracy(max_items: u64, fpp: f64) -> Self {
BloomFilterBuilder {
mode: BloomFilterBuilderMode::Accuracy { max_items, fpp },
seed: DEFAULT_UPDATE_SEED,
}
}
pub fn with_size(num_bits: u64, num_hashes: u16) -> Self {
BloomFilterBuilder {
mode: BloomFilterBuilderMode::Size {
num_bits,
num_hashes,
},
seed: DEFAULT_UPDATE_SEED,
}
}
pub fn seed(mut self, seed: u64) -> Self {
self.seed = seed;
self
}
pub fn build(self) -> Result<BloomFilter, Error> {
let (num_bits, num_hashes) = match self.mode {
BloomFilterBuilderMode::Accuracy { max_items, fpp } => {
if max_items == 0 {
return Err(Error::invalid_argument("max_items must be greater than 0"));
}
if !(fpp > 0.0 && fpp <= 1.0) {
return Err(Error::invalid_argument("fpp must be in (0.0, 1.0]"));
}
let n = max_items as f64;
let ln2_squared = std::f64::consts::LN_2 * std::f64::consts::LN_2;
let bits = (-n * fpp.ln() / ln2_squared).ceil();
if bits > Self::MAX_NUM_BITS as f64 {
return Err(Error::invalid_argument(format!(
"target accuracy requires {bits:.0} bits, but at most {} are supported",
Self::MAX_NUM_BITS
)));
}
let num_bits = (bits as u64).max(Self::MIN_NUM_BITS);
let num_hashes = (num_bits as f64 / n * std::f64::consts::LN_2).ceil().clamp(
f64::from(Self::MIN_NUM_HASHES),
f64::from(Self::MAX_NUM_HASHES),
) as u16;
(num_bits, num_hashes)
}
BloomFilterBuilderMode::Size {
num_bits,
num_hashes,
} => {
if !(Self::MIN_NUM_BITS..=Self::MAX_NUM_BITS).contains(&num_bits) {
return Err(Error::invalid_argument(format!(
"num_bits must be between {} and {}, got {}",
Self::MIN_NUM_BITS,
Self::MAX_NUM_BITS,
num_bits
)));
}
if !(Self::MIN_NUM_HASHES..=Self::MAX_NUM_HASHES).contains(&num_hashes) {
return Err(Error::invalid_argument(format!(
"num_hashes must be between {} and {}, got {}",
Self::MIN_NUM_HASHES,
Self::MAX_NUM_HASHES,
num_hashes
)));
}
(num_bits, num_hashes)
}
};
let num_words = num_bits.div_ceil(64) as usize;
let bit_array = vec![0u64; num_words].into_boxed_slice();
Ok(BloomFilter {
seed: self.seed,
num_hashes,
num_bits_set: 0,
bit_array,
})
}
}