use crate::error::{Result, TdbError};
use scirs2_core::metrics::{Counter, Gauge, Histogram, MetricsRegistry};
use scirs2_core::random::{rngs, Random, Rng};
use serde::{Deserialize, Serialize};
use std::hash::{Hash, Hasher};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BloomFilterConfig {
pub expected_elements: usize,
pub false_positive_rate: f64,
pub num_hash_functions: Option<usize>,
pub bit_array_size: Option<usize>,
pub enable_counting: bool,
pub enable_metrics: bool,
}
impl Default for BloomFilterConfig {
fn default() -> Self {
Self {
expected_elements: 10_000,
false_positive_rate: 0.01, num_hash_functions: None,
bit_array_size: None,
enable_counting: false,
enable_metrics: true,
}
}
}
impl BloomFilterConfig {
pub fn calculate_bit_array_size(&self) -> usize {
if let Some(size) = self.bit_array_size {
return size;
}
let n = self.expected_elements as f64;
let p = self.false_positive_rate;
let m = -(n * p.ln()) / (2_f64.ln().powi(2));
m.ceil() as usize
}
pub fn calculate_num_hash_functions(&self) -> usize {
if let Some(k) = self.num_hash_functions {
return k;
}
let m = self.calculate_bit_array_size() as f64;
let n = self.expected_elements as f64;
let k = (m / n) * 2_f64.ln();
k.ceil().max(1.0) as usize
}
}
pub struct BloomFilter {
config: BloomFilterConfig,
bits: Vec<AtomicU64>,
num_bits: usize,
num_hash_functions: usize,
hash_seeds: Vec<u64>,
rng: Random<rngs::StdRng>,
metrics: Option<BloomFilterMetrics>,
element_count: AtomicU64,
}
pub struct CountingBloomFilter {
config: BloomFilterConfig,
counters: Vec<AtomicU64>,
num_counters: usize,
num_hash_functions: usize,
hash_seeds: Vec<u64>,
rng: Random<rngs::StdRng>,
metrics: Option<BloomFilterMetrics>,
element_count: AtomicU64,
}
struct BloomFilterMetrics {
registry: Arc<MetricsRegistry>,
inserts: Counter,
deletes: Counter,
lookups: Counter,
true_positives: Counter,
false_positives: Counter,
fill_rate: Gauge,
lookup_latency: Histogram,
}
impl BloomFilterMetrics {
fn new(filter_name: &str) -> Self {
let registry = Arc::new(MetricsRegistry::new());
Self {
registry: registry.clone(),
inserts: Counter::new(format!("{}.inserts", filter_name)),
deletes: Counter::new(format!("{}.deletes", filter_name)),
lookups: Counter::new(format!("{}.lookups", filter_name)),
true_positives: Counter::new(format!("{}.true_positives", filter_name)),
false_positives: Counter::new(format!("{}.false_positives", filter_name)),
fill_rate: Gauge::new(format!("{}.fill_rate", filter_name)),
lookup_latency: Histogram::new(format!("{}.lookup_latency_us", filter_name)),
}
}
}
impl BloomFilter {
pub fn new(config: BloomFilterConfig) -> Result<Self> {
let num_bits = config.calculate_bit_array_size();
let num_hash_functions = config.calculate_num_hash_functions();
let num_u64s = (num_bits + 63) / 64;
let mut rng = Random::seed(0);
let hash_seeds: Vec<u64> = (0..num_hash_functions).map(|_| rng.next_u64()).collect();
let metrics = if config.enable_metrics {
Some(BloomFilterMetrics::new("bloom_filter"))
} else {
None
};
Ok(Self {
config,
bits: (0..num_u64s).map(|_| AtomicU64::new(0)).collect(),
num_bits,
num_hash_functions,
hash_seeds,
rng,
metrics,
element_count: AtomicU64::new(0),
})
}
pub fn insert<T: Hash>(&mut self, item: &T) {
let start = std::time::Instant::now();
for &seed in &self.hash_seeds {
let hash = self.hash_with_seed(item, seed);
let bit_index = (hash % self.num_bits as u64) as usize;
let word_index = bit_index / 64;
let bit_offset = bit_index % 64;
self.bits[word_index].fetch_or(1u64 << bit_offset, Ordering::Relaxed);
}
self.element_count.fetch_add(1, Ordering::Relaxed);
if let Some(ref metrics) = self.metrics {
metrics.inserts.inc();
self.update_fill_rate();
}
}
pub fn contains<T: Hash>(&self, item: &T) -> bool {
let start = std::time::Instant::now();
let result = self.hash_seeds.iter().all(|&seed| {
let hash = self.hash_with_seed(item, seed);
let bit_index = (hash % self.num_bits as u64) as usize;
let word_index = bit_index / 64;
let bit_offset = bit_index % 64;
let word = self.bits[word_index].load(Ordering::Relaxed);
(word & (1u64 << bit_offset)) != 0
});
if let Some(ref metrics) = self.metrics {
metrics.lookups.inc();
let elapsed = start.elapsed();
metrics.lookup_latency.observe(elapsed.as_micros() as f64);
}
result
}
pub fn clear(&mut self) {
for word in &self.bits {
word.store(0, Ordering::Relaxed);
}
self.element_count.store(0, Ordering::Relaxed);
}
pub fn estimated_false_positive_rate(&self) -> f64 {
let n = self.element_count.load(Ordering::Relaxed) as f64;
let m = self.num_bits as f64;
let k = self.num_hash_functions as f64;
(1.0 - (-k * n / m).exp()).powf(k)
}
pub fn fill_rate(&self) -> f64 {
let total_bits = self.num_bits;
let set_bits: u64 = self
.bits
.iter()
.map(|word| word.load(Ordering::Relaxed).count_ones() as u64)
.sum();
set_bits as f64 / total_bits as f64
}
pub fn stats(&self) -> BloomFilterStats {
BloomFilterStats {
num_bits: self.num_bits,
num_hash_functions: self.num_hash_functions,
element_count: self.element_count.load(Ordering::Relaxed),
fill_rate: self.fill_rate(),
estimated_fpr: self.estimated_false_positive_rate(),
configured_fpr: self.config.false_positive_rate,
}
}
fn hash_with_seed<T: Hash>(&self, item: &T, seed: u64) -> u64 {
let mut hasher = std::collections::hash_map::DefaultHasher::new();
seed.hash(&mut hasher);
item.hash(&mut hasher);
hasher.finish()
}
fn update_fill_rate(&self) {
if let Some(ref metrics) = self.metrics {
metrics.fill_rate.set(self.fill_rate());
}
}
}
impl CountingBloomFilter {
pub fn new(config: BloomFilterConfig) -> Result<Self> {
let num_counters = config.calculate_bit_array_size();
let num_hash_functions = config.calculate_num_hash_functions();
let mut rng = Random::seed(0);
let hash_seeds: Vec<u64> = (0..num_hash_functions).map(|_| rng.next_u64()).collect();
let metrics = if config.enable_metrics {
Some(BloomFilterMetrics::new("counting_bloom_filter"))
} else {
None
};
Ok(Self {
config,
counters: (0..num_counters).map(|_| AtomicU64::new(0)).collect(),
num_counters,
num_hash_functions,
hash_seeds,
rng,
metrics,
element_count: AtomicU64::new(0),
})
}
pub fn insert<T: Hash>(&mut self, item: &T) {
for &seed in &self.hash_seeds {
let hash = self.hash_with_seed(item, seed);
let counter_index = (hash % self.num_counters as u64) as usize;
let current = self.counters[counter_index].load(Ordering::Relaxed);
if current < 15 {
self.counters[counter_index].fetch_add(1, Ordering::Relaxed);
}
}
self.element_count.fetch_add(1, Ordering::Relaxed);
if let Some(ref metrics) = self.metrics {
metrics.inserts.inc();
}
}
pub fn delete<T: Hash>(&mut self, item: &T) {
for &seed in &self.hash_seeds {
let hash = self.hash_with_seed(item, seed);
let counter_index = (hash % self.num_counters as u64) as usize;
let current = self.counters[counter_index].load(Ordering::Relaxed);
if current > 0 {
self.counters[counter_index].fetch_sub(1, Ordering::Relaxed);
}
}
self.element_count.fetch_sub(1, Ordering::Relaxed);
if let Some(ref metrics) = self.metrics {
metrics.deletes.inc();
}
}
pub fn contains<T: Hash>(&self, item: &T) -> bool {
let start = std::time::Instant::now();
let result = self.hash_seeds.iter().all(|&seed| {
let hash = self.hash_with_seed(item, seed);
let counter_index = (hash % self.num_counters as u64) as usize;
self.counters[counter_index].load(Ordering::Relaxed) > 0
});
if let Some(ref metrics) = self.metrics {
metrics.lookups.inc();
let elapsed = start.elapsed();
metrics.lookup_latency.observe(elapsed.as_micros() as f64);
}
result
}
pub fn stats(&self) -> BloomFilterStats {
let total_counters = self.num_counters;
let non_zero_counters = self
.counters
.iter()
.filter(|c| c.load(Ordering::Relaxed) > 0)
.count();
let element_count = self.element_count.load(Ordering::Relaxed);
let estimated_fpr = if element_count > 0 {
let k = self.num_hash_functions as f64;
let n = element_count as f64;
let m = self.num_counters as f64;
let exponent = -k * n / m;
let base = 1.0 - exponent.exp();
base.powf(k)
} else {
0.0
};
BloomFilterStats {
num_bits: self.num_counters, num_hash_functions: self.num_hash_functions,
element_count,
fill_rate: non_zero_counters as f64 / total_counters as f64,
estimated_fpr,
configured_fpr: self.config.false_positive_rate,
}
}
fn hash_with_seed<T: Hash>(&self, item: &T, seed: u64) -> u64 {
let mut hasher = std::collections::hash_map::DefaultHasher::new();
seed.hash(&mut hasher);
item.hash(&mut hasher);
hasher.finish()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BloomFilterStats {
pub num_bits: usize,
pub num_hash_functions: usize,
pub element_count: u64,
pub fill_rate: f64,
pub estimated_fpr: f64,
pub configured_fpr: f64,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_bloom_filter_basic() -> Result<()> {
let config = BloomFilterConfig {
expected_elements: 100,
false_positive_rate: 0.01,
..Default::default()
};
let mut filter = BloomFilter::new(config)?;
filter.insert(&"hello");
filter.insert(&"world");
filter.insert(&123u64);
assert!(filter.contains(&"hello"));
assert!(filter.contains(&"world"));
assert!(filter.contains(&123u64));
assert!(!filter.contains(&"not_present"));
Ok(())
}
#[test]
fn test_bloom_filter_stats() -> Result<()> {
let config = BloomFilterConfig::default();
let mut filter = BloomFilter::new(config)?;
for i in 0..100 {
filter.insert(&i);
}
let stats = filter.stats();
assert_eq!(stats.element_count, 100);
assert!(stats.fill_rate > 0.0);
assert!(stats.fill_rate < 1.0);
Ok(())
}
#[test]
fn test_counting_bloom_filter() -> Result<()> {
let config = BloomFilterConfig {
expected_elements: 100,
enable_counting: true,
..Default::default()
};
let mut filter = CountingBloomFilter::new(config)?;
filter.insert(&"test");
assert!(filter.contains(&"test"));
filter.delete(&"test");
assert!(!filter.contains(&"test"));
Ok(())
}
#[test]
fn test_false_positive_rate() -> Result<()> {
let config = BloomFilterConfig {
expected_elements: 1000,
false_positive_rate: 0.01,
..Default::default()
};
let mut filter = BloomFilter::new(config)?;
for i in 0..1000 {
filter.insert(&i);
}
let mut false_positives = 0;
for i in 1000..2000 {
if filter.contains(&i) {
false_positives += 1;
}
}
let actual_fpr = false_positives as f64 / 1000.0;
assert!(actual_fpr < 0.05);
Ok(())
}
#[test]
fn test_clear() -> Result<()> {
let config = BloomFilterConfig::default();
let mut filter = BloomFilter::new(config)?;
filter.insert(&"test");
assert!(filter.contains(&"test"));
filter.clear();
assert!(!filter.contains(&"test"));
let stats = filter.stats();
assert_eq!(stats.element_count, 0);
assert_eq!(stats.fill_rate, 0.0);
Ok(())
}
#[test]
fn test_config_calculations() {
let config = BloomFilterConfig {
expected_elements: 10_000,
false_positive_rate: 0.01,
..Default::default()
};
let bits = config.calculate_bit_array_size();
let hash_funcs = config.calculate_num_hash_functions();
assert!(bits > 0);
assert!(hash_funcs > 0);
assert!(hash_funcs < 20); }
#[test]
fn test_multiple_insertions() -> Result<()> {
let config = BloomFilterConfig::default();
let mut filter = BloomFilter::new(config)?;
filter.insert(&"test");
filter.insert(&"test");
filter.insert(&"test");
assert!(filter.contains(&"test"));
Ok(())
}
#[test]
fn test_counting_filter_multiple_ops() -> Result<()> {
let config = BloomFilterConfig {
enable_counting: true,
..Default::default()
};
let mut filter = CountingBloomFilter::new(config)?;
filter.insert(&"key");
filter.insert(&"key");
assert!(filter.contains(&"key"));
filter.delete(&"key");
assert!(filter.contains(&"key"));
filter.delete(&"key");
assert!(!filter.contains(&"key"));
Ok(())
}
}