use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
#[derive(Clone, Debug)]
pub struct BloomFilter {
bits: Vec<u8>,
num_hashes: u32,
num_bits: usize,
}
impl BloomFilter {
pub fn new(num_keys: usize, bits_per_key: usize) -> Self {
let num_bits = (num_keys * bits_per_key).max(64); let num_bytes = num_bits.div_ceil(8);
let num_hashes = ((bits_per_key as f64) * 0.693).ceil() as u32;
let num_hashes = num_hashes.clamp(1, 30);
Self {
bits: vec![0u8; num_bytes],
num_hashes,
num_bits,
}
}
pub fn from_bytes(bits: Vec<u8>, num_hashes: u32) -> Self {
let num_bits = bits.len() * 8;
Self {
bits,
num_hashes,
num_bits,
}
}
pub fn insert(&mut self, key: &[u8]) {
if self.num_bits == 0 {
return;
}
for i in 0..self.num_hashes {
let hash = self.hash(key, i);
let bit_pos = (hash as usize) % self.num_bits;
self.set_bit(bit_pos);
}
}
pub fn may_contain(&self, key: &[u8]) -> bool {
if self.num_bits == 0 {
return false;
}
for i in 0..self.num_hashes {
let hash = self.hash(key, i);
let bit_pos = (hash as usize) % self.num_bits;
if !self.get_bit(bit_pos) {
return false; }
}
true }
pub fn may_contain_batch(&self, keys: &[&[u8]]) -> Vec<bool> {
if self.num_bits == 0 {
return vec![false; keys.len()];
}
let mut results = vec![false; keys.len()];
let mut hash_cache: Vec<Vec<u64>> = Vec::with_capacity(keys.len());
for key in keys {
let mut hashes = Vec::with_capacity(self.num_hashes as usize);
for i in 0..self.num_hashes {
let hash = self.hash(key, i);
hashes.push(hash);
}
hash_cache.push(hashes);
}
for (idx, hashes) in hash_cache.iter().enumerate() {
let mut found = true;
for &hash in hashes {
let bit_pos = (hash as usize) % self.num_bits;
if !self.get_bit(bit_pos) {
found = false;
break; }
}
results[idx] = found;
}
results
}
pub fn to_bytes(&self) -> Vec<u8> {
let mut buf = Vec::new();
buf.extend_from_slice(&self.num_hashes.to_le_bytes());
buf.extend_from_slice(&(self.num_bits as u64).to_le_bytes());
buf.extend_from_slice(&self.bits);
buf
}
pub fn from_bytes_full(data: &[u8]) -> Option<Self> {
if data.len() < 12 {
return None;
}
let num_hashes = u32::from_le_bytes([data[0], data[1], data[2], data[3]]);
let num_bits = u64::from_le_bytes([
data[4], data[5], data[6], data[7], data[8], data[9], data[10], data[11],
]) as usize;
let bits = data[12..].to_vec();
if num_bits > bits.len() * 8 || num_hashes == 0 || num_hashes > 30 {
return None;
}
Some(Self {
bits,
num_hashes,
num_bits,
})
}
pub fn byte_size(&self) -> usize {
12 + self.bits.len() }
fn hash(&self, key: &[u8], seed: u32) -> u64 {
let mut hasher = DefaultHasher::new();
seed.hash(&mut hasher);
key.hash(&mut hasher);
hasher.finish()
}
fn set_bit(&mut self, pos: usize) {
let byte_idx = pos / 8;
let bit_idx = pos % 8;
self.bits[byte_idx] |= 1 << bit_idx;
}
fn get_bit(&self, pos: usize) -> bool {
let byte_idx = pos / 8;
let bit_idx = pos % 8;
(self.bits[byte_idx] & (1 << bit_idx)) != 0
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_operations() {
let mut bloom = BloomFilter::new(100, 10);
bloom.insert(b"key1");
bloom.insert(b"key2");
bloom.insert(b"key3");
assert!(bloom.may_contain(b"key1"));
assert!(bloom.may_contain(b"key2"));
assert!(bloom.may_contain(b"key3"));
assert!(!bloom.may_contain(b"key4"));
assert!(!bloom.may_contain(b"key5"));
}
#[test]
fn test_false_positive_rate() {
let num_keys = 1000;
let bits_per_key = 10;
let mut bloom = BloomFilter::new(num_keys, bits_per_key);
for i in 0..num_keys {
let key = format!("key_{}", i);
bloom.insert(key.as_bytes());
}
let mut false_positives = 0;
let test_count = 10000;
for i in num_keys..(num_keys + test_count) {
let key = format!("key_{}", i);
if bloom.may_contain(key.as_bytes()) {
false_positives += 1;
}
}
let fpr = false_positives as f64 / test_count as f64;
debug_log!("False positive rate: {:.2}%", fpr * 100.0);
assert!(fpr < 0.03, "FPR too high: {:.2}%", fpr * 100.0);
}
#[test]
fn test_serialization() {
let mut bloom = BloomFilter::new(100, 10);
bloom.insert(b"key1");
bloom.insert(b"key2");
let bytes = bloom.to_bytes();
let bloom2 = BloomFilter::from_bytes_full(&bytes).unwrap();
assert!(bloom2.may_contain(b"key1"));
assert!(bloom2.may_contain(b"key2"));
assert!(!bloom2.may_contain(b"nonexistent"));
}
#[test]
fn test_empty_filter() {
let bloom = BloomFilter::new(100, 10); assert!(!bloom.may_contain(b"any_key"));
}
}