use borsh::{BorshDeserialize, BorshSerialize};
use super::hash_comparison::TreeLeafData;
pub const DEFAULT_BLOOM_FP_RATE: f64 = 0.01;
const MIN_BITS_PER_ELEMENT: usize = 8;
const MIN_FP_RATE: f64 = 0.0001;
const MAX_FP_RATE: f64 = 0.5;
const MIN_NUM_BITS: usize = 64;
const MIN_NUM_HASHES: u8 = 1;
const MAX_NUM_HASHES: u8 = 16;
const FNV_OFFSET_BASIS: u64 = 0xcbf29ce484222325;
const FNV_PRIME: u64 = 0x100000001b3;
#[derive(Clone, Debug, PartialEq, BorshSerialize, BorshDeserialize)]
pub struct DeltaIdBloomFilter {
bits: Vec<u8>,
num_bits: usize,
num_hashes: u8,
item_count: usize,
}
impl DeltaIdBloomFilter {
#[must_use]
pub fn new(expected_items: usize, fp_rate: f64) -> Self {
let fp_rate = Self::clamp_fp_rate(fp_rate);
let ln2_sq = std::f64::consts::LN_2 * std::f64::consts::LN_2;
let num_bits = if expected_items == 0 {
MIN_NUM_BITS
} else {
let m = -(expected_items as f64) * fp_rate.ln() / ln2_sq;
(m.ceil() as usize).max(expected_items.saturating_mul(MIN_BITS_PER_ELEMENT))
};
let num_hashes = if expected_items == 0 {
4 } else {
let k = (num_bits as f64 / expected_items as f64) * std::f64::consts::LN_2;
(k.ceil() as u8).clamp(MIN_NUM_HASHES, MAX_NUM_HASHES)
};
let num_bytes = (num_bits + 7) / 8;
Self {
bits: vec![0; num_bytes],
num_bits,
num_hashes,
item_count: 0,
}
}
#[must_use]
pub fn with_params(num_bits: usize, num_hashes: u8) -> Self {
let num_bits = num_bits.max(MIN_NUM_BITS);
let num_hashes = num_hashes.clamp(MIN_NUM_HASHES, MAX_NUM_HASHES);
let num_bytes = (num_bits + 7) / 8;
Self {
bits: vec![0; num_bytes],
num_bits,
num_hashes,
item_count: 0,
}
}
#[must_use]
pub fn hash_fnv1a(data: &[u8]) -> u64 {
let mut hash: u64 = FNV_OFFSET_BASIS;
for byte in data {
hash ^= *byte as u64;
hash = hash.wrapping_mul(FNV_PRIME);
}
hash
}
#[inline]
fn compute_hashes(id: &[u8; 32]) -> (u64, u64) {
let h1 = Self::hash_fnv1a(id);
let mut buf = [0u8; 33];
buf[..32].copy_from_slice(id);
buf[32] = 0xFF;
let h2 = Self::hash_fnv1a(&buf);
(h1, h2)
}
#[inline]
fn position_at(&self, h1: u64, h2: u64, i: u64) -> usize {
let combined = h1.wrapping_add(i.wrapping_mul(h2));
(combined as usize) % self.num_bits
}
pub fn insert(&mut self, id: &[u8; 32]) {
if !self.is_valid() {
return;
}
let (h1, h2) = Self::compute_hashes(id);
for i in 0..self.num_hashes as u64 {
let pos = self.position_at(h1, h2, i);
let byte_idx = pos / 8;
let bit_idx = pos % 8;
if byte_idx < self.bits.len() {
self.bits[byte_idx] |= 1 << bit_idx;
}
}
self.item_count += 1;
}
#[must_use]
pub fn contains(&self, id: &[u8; 32]) -> bool {
if !self.is_valid() {
return false;
}
let (h1, h2) = Self::compute_hashes(id);
for i in 0..self.num_hashes as u64 {
let pos = self.position_at(h1, h2, i);
let byte_idx = pos / 8;
let bit_idx = pos % 8;
if byte_idx >= self.bits.len() {
return false;
}
if self.bits[byte_idx] & (1 << bit_idx) == 0 {
return false;
}
}
true
}
#[must_use]
pub fn item_count(&self) -> usize {
self.item_count
}
#[must_use]
pub fn bit_count(&self) -> usize {
self.num_bits
}
#[must_use]
pub fn hash_count(&self) -> u8 {
self.num_hashes
}
#[must_use]
pub fn estimated_fp_rate(&self) -> f64 {
if self.item_count == 0 {
return 0.0;
}
let k = self.num_hashes as f64;
let n = self.item_count as f64;
let m = self.num_bits as f64;
(1.0 - (-k * n / m).exp()).powf(k)
}
#[must_use]
pub fn bits(&self) -> &[u8] {
&self.bits
}
#[must_use]
pub fn is_valid(&self) -> bool {
if self.num_bits == 0 || self.num_hashes == 0 {
return false;
}
self.num_bits <= self.bits.len().saturating_mul(8)
}
#[must_use]
pub fn clamp_fp_rate(fp_rate: f64) -> f64 {
if fp_rate.is_nan() {
DEFAULT_BLOOM_FP_RATE
} else {
fp_rate.clamp(MIN_FP_RATE, MAX_FP_RATE)
}
}
}
#[derive(Clone, Debug, PartialEq, BorshSerialize, BorshDeserialize)]
pub struct BloomFilterRequest {
pub filter: DeltaIdBloomFilter,
pub false_positive_rate: f64,
}
impl BloomFilterRequest {
#[must_use]
pub fn new(filter: DeltaIdBloomFilter, false_positive_rate: f64) -> Self {
Self {
filter,
false_positive_rate,
}
}
#[must_use]
pub fn from_ids(ids: &[[u8; 32]], fp_rate: f64) -> Self {
let clamped_fp_rate = DeltaIdBloomFilter::clamp_fp_rate(fp_rate);
let mut filter = DeltaIdBloomFilter::new(ids.len(), clamped_fp_rate);
for id in ids {
filter.insert(id);
}
Self::new(filter, clamped_fp_rate)
}
}
#[derive(Clone, Debug, PartialEq, BorshSerialize, BorshDeserialize)]
pub struct BloomFilterResponse {
pub missing_entities: Vec<TreeLeafData>,
pub scanned_count: usize,
}
impl BloomFilterResponse {
#[must_use]
pub fn new(missing_entities: Vec<TreeLeafData>, scanned_count: usize) -> Self {
Self {
missing_entities,
scanned_count,
}
}
#[must_use]
pub fn empty(scanned_count: usize) -> Self {
Self {
missing_entities: vec![],
scanned_count,
}
}
#[must_use]
pub fn has_missing(&self) -> bool {
!self.missing_entities.is_empty()
}
#[must_use]
pub fn missing_count(&self) -> usize {
self.missing_entities.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sync::hash_comparison::{CrdtType, LeafMetadata};
#[test]
fn test_bloom_filter_fnv1a_consistency() {
let data = [1u8; 32];
let hash1 = DeltaIdBloomFilter::hash_fnv1a(&data);
let hash2 = DeltaIdBloomFilter::hash_fnv1a(&data);
assert_eq!(hash1, hash2);
let other_data = [2u8; 32];
let other_hash = DeltaIdBloomFilter::hash_fnv1a(&other_data);
assert_ne!(hash1, other_hash);
}
#[test]
fn test_bloom_filter_insert_contains() {
let mut filter = DeltaIdBloomFilter::new(100, 0.01);
let id1 = [1u8; 32];
let id2 = [2u8; 32];
let id3 = [3u8; 32];
assert!(!filter.contains(&id1));
assert!(!filter.contains(&id2));
filter.insert(&id1);
filter.insert(&id2);
assert!(filter.contains(&id1));
assert!(filter.contains(&id2));
assert!(!filter.contains(&id3));
}
#[test]
fn test_bloom_filter_item_count() {
let mut filter = DeltaIdBloomFilter::new(100, 0.01);
assert_eq!(filter.item_count(), 0);
filter.insert(&[1u8; 32]);
assert_eq!(filter.item_count(), 1);
filter.insert(&[2u8; 32]);
filter.insert(&[3u8; 32]);
assert_eq!(filter.item_count(), 3);
}
#[test]
fn test_bloom_filter_roundtrip() {
let mut filter = DeltaIdBloomFilter::new(50, 0.01);
filter.insert(&[1u8; 32]);
filter.insert(&[2u8; 32]);
filter.insert(&[3u8; 32]);
let encoded = borsh::to_vec(&filter).expect("serialize");
let decoded: DeltaIdBloomFilter = borsh::from_slice(&encoded).expect("deserialize");
assert_eq!(filter, decoded);
assert!(decoded.contains(&[1u8; 32]));
assert!(decoded.contains(&[2u8; 32]));
assert!(decoded.contains(&[3u8; 32]));
assert!(!decoded.contains(&[4u8; 32]));
}
#[test]
fn test_bloom_filter_false_positive_rate() {
let num_items = 1000;
let target_fp_rate = 0.01;
let mut filter = DeltaIdBloomFilter::new(num_items, target_fp_rate);
for i in 0..num_items {
let mut id = [0u8; 32];
id[0..8].copy_from_slice(&(i as u64).to_le_bytes());
filter.insert(&id);
}
let test_count = 10000;
let mut false_positives = 0;
for i in num_items..(num_items + test_count) {
let mut id = [0u8; 32];
id[0..8].copy_from_slice(&(i as u64).to_le_bytes());
if filter.contains(&id) {
false_positives += 1;
}
}
let actual_fp_rate = false_positives as f64 / test_count as f64;
assert!(
actual_fp_rate < target_fp_rate as f64 * 3.0,
"FP rate {} too high (target {})",
actual_fp_rate,
target_fp_rate
);
}
#[test]
fn test_bloom_filter_estimated_fp_rate() {
let mut filter = DeltaIdBloomFilter::new(100, 0.01);
assert_eq!(filter.estimated_fp_rate(), 0.0);
for i in 0..50 {
let mut id = [0u8; 32];
id[0..8].copy_from_slice(&(i as u64).to_le_bytes());
filter.insert(&id);
}
let estimated = filter.estimated_fp_rate();
assert!(estimated > 0.0);
assert!(estimated < 0.1);
}
#[test]
fn test_bloom_filter_request_from_ids() {
let ids = [[1u8; 32], [2u8; 32], [3u8; 32]];
let request = BloomFilterRequest::from_ids(&ids, 0.01);
assert!(request.filter.contains(&[1u8; 32]));
assert!(request.filter.contains(&[2u8; 32]));
assert!(request.filter.contains(&[3u8; 32]));
assert!(!request.filter.contains(&[4u8; 32]));
assert_eq!(request.false_positive_rate, 0.01);
}
#[test]
fn test_bloom_filter_request_roundtrip() {
let ids = [[1u8; 32], [2u8; 32]];
let request = BloomFilterRequest::from_ids(&ids, 0.02);
let encoded = borsh::to_vec(&request).expect("serialize");
let decoded: BloomFilterRequest = borsh::from_slice(&encoded).expect("deserialize");
assert_eq!(request, decoded);
}
#[test]
fn test_bloom_filter_response() {
let metadata = LeafMetadata::new(CrdtType::lww_register("test"), 100, [5; 32]);
let leaf = TreeLeafData::new([1; 32], vec![1, 2, 3], metadata);
let response = BloomFilterResponse::new(vec![leaf], 100);
assert!(response.has_missing());
assert_eq!(response.missing_count(), 1);
assert_eq!(response.scanned_count, 100);
}
#[test]
fn test_bloom_filter_response_empty() {
let response = BloomFilterResponse::empty(50);
assert!(!response.has_missing());
assert_eq!(response.missing_count(), 0);
assert_eq!(response.scanned_count, 50);
}
#[test]
fn test_bloom_filter_response_roundtrip() {
let metadata = LeafMetadata::new(CrdtType::unordered_map("String", "u64"), 200, [6; 32]);
let leaf = TreeLeafData::new([2; 32], vec![4, 5, 6], metadata);
let response = BloomFilterResponse::new(vec![leaf], 75);
let encoded = borsh::to_vec(&response).expect("serialize");
let decoded: BloomFilterResponse = borsh::from_slice(&encoded).expect("deserialize");
assert_eq!(response, decoded);
}
#[test]
fn test_bloom_filter_with_params() {
let filter = DeltaIdBloomFilter::with_params(1024, 7);
assert_eq!(filter.bit_count(), 1024);
assert_eq!(filter.hash_count(), 7);
assert_eq!(filter.item_count(), 0);
}
#[test]
fn test_bloom_filter_fp_rate_clamping_zero() {
let filter = DeltaIdBloomFilter::new(100, 0.0);
assert!(filter.bit_count() > 0);
assert!(filter.hash_count() > 0);
let id = [42u8; 32];
assert!(!filter.contains(&id));
}
#[test]
fn test_bloom_filter_fp_rate_clamping_negative() {
let filter = DeltaIdBloomFilter::new(100, -0.5);
assert!(filter.bit_count() > 0);
assert!(filter.hash_count() > 0);
}
#[test]
fn test_bloom_filter_fp_rate_clamping_too_high() {
let filter = DeltaIdBloomFilter::new(100, 0.99);
assert!(filter.bit_count() > 0);
assert!(filter.hash_count() > 0);
let mut filter = filter;
let id = [1u8; 32];
filter.insert(&id);
assert!(filter.contains(&id));
}
#[test]
fn test_bloom_filter_fp_rate_edge_cases() {
let test_cases = [
(f64::NEG_INFINITY, "negative infinity"),
(f64::INFINITY, "positive infinity"),
(f64::NAN, "NaN"),
(-1.0_f64, "negative one"),
(0.0_f64, "zero"),
(1.0_f64, "one"),
(2.0_f64, "greater than one"),
];
for (fp_rate, description) in test_cases {
let filter = DeltaIdBloomFilter::new(100, fp_rate);
assert!(
filter.bit_count() > 0,
"Filter with fp_rate {} ({}) should have positive bit count",
fp_rate,
description
);
}
}
#[test]
fn test_bloom_filter_nan_uses_default() {
let nan_filter = DeltaIdBloomFilter::new(100, f64::NAN);
let default_filter = DeltaIdBloomFilter::new(100, DEFAULT_BLOOM_FP_RATE);
assert_eq!(nan_filter.bit_count(), default_filter.bit_count());
assert_eq!(nan_filter.hash_count(), default_filter.hash_count());
}
#[test]
fn test_bloom_filter_hashing_deterministic() {
let id = [0xAB; 32];
let (h1_a, h2_a) = DeltaIdBloomFilter::compute_hashes(&id);
let (h1_b, h2_b) = DeltaIdBloomFilter::compute_hashes(&id);
assert_eq!(h1_a, h1_b);
assert_eq!(h2_a, h2_b);
let other_id = [0xCD; 32];
let (h1_other, h2_other) = DeltaIdBloomFilter::compute_hashes(&other_id);
assert_ne!(h1_a, h1_other);
}
#[test]
fn test_bloom_filter_insert_contains_deterministic() {
let mut filter1 = DeltaIdBloomFilter::new(100, 0.01);
let mut filter2 = DeltaIdBloomFilter::new(100, 0.01);
let id = [0xAB; 32];
filter1.insert(&id);
filter2.insert(&id);
assert_eq!(filter1.bits(), filter2.bits());
assert!(filter1.contains(&id));
assert!(filter2.contains(&id));
}
#[test]
fn test_bloom_filter_with_params_clamps_num_bits() {
let filter = DeltaIdBloomFilter::with_params(0, 4);
assert_eq!(filter.bit_count(), 64);
assert!(filter.is_valid());
let id = [1u8; 32];
assert!(!filter.contains(&id));
}
#[test]
fn test_bloom_filter_with_params_clamps_num_hashes() {
let filter = DeltaIdBloomFilter::with_params(128, 0);
assert_eq!(filter.hash_count(), 1);
assert!(filter.is_valid());
let filter_high = DeltaIdBloomFilter::with_params(128, 255);
assert_eq!(filter_high.hash_count(), 16);
}
#[test]
fn test_bloom_filter_is_valid() {
let valid_new = DeltaIdBloomFilter::new(100, 0.01);
assert!(valid_new.is_valid());
let valid_params = DeltaIdBloomFilter::with_params(0, 0);
assert!(valid_params.is_valid()); }
#[test]
fn test_bloom_filter_malicious_deserialization_num_bits_zero() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&0usize.to_le_bytes()); bytes.push(4);
bytes.extend_from_slice(&0usize.to_le_bytes());
let filter: DeltaIdBloomFilter = borsh::from_slice(&bytes).expect("deserialize");
assert!(
!filter.is_valid(),
"Filter with num_bits=0 should be invalid"
);
let id = [1u8; 32];
assert!(!filter.contains(&id));
}
#[test]
fn test_bloom_filter_malicious_deserialization_num_hashes_zero() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&8u32.to_le_bytes()); bytes.extend_from_slice(&[0u8; 8]); bytes.extend_from_slice(&64usize.to_le_bytes());
bytes.push(0); bytes.extend_from_slice(&0usize.to_le_bytes());
let filter: DeltaIdBloomFilter = borsh::from_slice(&bytes).expect("deserialize");
assert!(
!filter.is_valid(),
"Filter with num_hashes=0 should be invalid"
);
}
#[test]
fn test_bloom_filter_malicious_deserialization_bits_too_small() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&1u32.to_le_bytes()); bytes.push(0u8); bytes.extend_from_slice(&1_000_000usize.to_le_bytes());
bytes.push(4);
bytes.extend_from_slice(&0usize.to_le_bytes());
let filter: DeltaIdBloomFilter = borsh::from_slice(&bytes).expect("deserialize");
assert!(
!filter.is_valid(),
"Filter with bits.len()=1 but num_bits=1000000 should be invalid"
);
let id = [0xAB; 32];
assert!(!filter.contains(&id));
let mut filter_mut = filter;
filter_mut.insert(&id); }
#[test]
fn test_bloom_filter_structural_consistency() {
let filter = DeltaIdBloomFilter::new(100, 0.01);
assert!(filter.is_valid());
let required_bytes = (filter.bit_count() + 7) / 8;
assert!(filter.bits().len() >= required_bytes);
}
#[test]
fn test_bloom_filter_malicious_deserialization_num_bits_max() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&0u32.to_le_bytes()); bytes.extend_from_slice(&usize::MAX.to_le_bytes());
bytes.push(4);
bytes.extend_from_slice(&0usize.to_le_bytes());
let filter: DeltaIdBloomFilter = borsh::from_slice(&bytes).expect("deserialize");
assert!(
!filter.is_valid(),
"Filter with num_bits=usize::MAX but empty bits should be invalid"
);
let id = [0xAB; 32];
assert!(!filter.contains(&id));
}
#[test]
fn test_bloom_filter_from_ids_clamps_fp_rate() {
let ids = [[1u8; 32], [2u8; 32]];
let request = BloomFilterRequest::from_ids(&ids, f64::NAN);
assert_eq!(
request.false_positive_rate, DEFAULT_BLOOM_FP_RATE,
"NaN should be clamped to DEFAULT_BLOOM_FP_RATE"
);
let request = BloomFilterRequest::from_ids(&ids, -1.0);
assert_eq!(
request.false_positive_rate, MIN_FP_RATE,
"Negative should be clamped to MIN_FP_RATE"
);
let request = BloomFilterRequest::from_ids(&ids, 0.99);
assert_eq!(
request.false_positive_rate, MAX_FP_RATE,
"Value > 0.5 should be clamped to MAX_FP_RATE"
);
let request = BloomFilterRequest::from_ids(&ids, 0.02);
assert_eq!(
request.false_positive_rate, 0.02,
"Valid FP rate should remain unchanged"
);
}
#[test]
fn test_clamp_fp_rate() {
assert_eq!(
DeltaIdBloomFilter::clamp_fp_rate(f64::NAN),
DEFAULT_BLOOM_FP_RATE
);
assert_eq!(DeltaIdBloomFilter::clamp_fp_rate(-1.0), MIN_FP_RATE);
assert_eq!(DeltaIdBloomFilter::clamp_fp_rate(0.0), MIN_FP_RATE);
assert_eq!(DeltaIdBloomFilter::clamp_fp_rate(0.99), MAX_FP_RATE);
assert_eq!(
DeltaIdBloomFilter::clamp_fp_rate(f64::INFINITY),
MAX_FP_RATE
);
assert_eq!(
DeltaIdBloomFilter::clamp_fp_rate(f64::NEG_INFINITY),
MIN_FP_RATE
);
assert_eq!(DeltaIdBloomFilter::clamp_fp_rate(0.01), 0.01);
assert_eq!(DeltaIdBloomFilter::clamp_fp_rate(0.1), 0.1);
}
}