use crate::banding::{Banding, Solution, W};
use crate::hash::{ribbon_hash, start};
use crate::PleatPlan;
pub const DEFAULT_WINDOW_SHIFT: u32 = 16;
pub(crate) fn window_size(window_shift: u32) -> usize {
assert!(
window_shift < usize::BITS,
"pleating window shift must be smaller than the target pointer width"
);
1usize << window_shift
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BuildError {
NoSolvingSeed,
}
impl core::fmt::Display for BuildError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::NoSolvingSeed => f.write_str("no standard ribbon seed solved the key set"),
}
}
}
impl std::error::Error for BuildError {}
fn overhead(r: usize) -> f64 {
1.0 + (4.0 + r as f64 * 0.25) / (8.0 * 8.0)
}
pub(crate) fn num_slots_for(n: usize, r: usize) -> usize {
let raw = (overhead(r) * n as f64) as usize;
let mut s = raw.div_ceil(W) * W;
if s == W {
s += W;
}
s.max(2 * W)
}
pub struct Ribbon<const R: usize> {
soln: Solution<R>,
}
pub type RibbonFilter = Ribbon<7>;
impl<const R: usize> core::fmt::Debug for Ribbon<R> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("Ribbon")
.field("w", &64u32)
.field("r", &R)
.field("bytes", &self.size_bytes())
.finish()
}
}
impl<const R: usize> Ribbon<R> {
pub fn from_keys(keys: &[u64]) -> Self {
Self::from_keys_seeded(keys, 0)
}
pub fn from_keys_seeded(keys: &[u64], seed: u64) -> Self {
Self::check_width();
let mut b = Banding::<R>::new(num_slots_for(keys.len(), R), seed);
b.add_all(keys);
Self { soln: b.solve() }
}
pub fn from_keys_pleated(keys: &[u64]) -> Self {
Self::from_keys_pleated_seeded(keys, 0, DEFAULT_WINDOW_SHIFT)
}
pub fn from_keys_pleated_seeded(keys: &[u64], seed: u64, window_shift: u32) -> Self {
Self::check_width();
let _ = window_size(window_shift);
let num_slots = num_slots_for(keys.len(), R);
let num_starts = (num_slots - W + 1) as u64;
let plan = PleatPlan::new(num_starts, window_shift);
let (ordered, _counts) = plan.pleat(keys, |k| start(ribbon_hash(k, seed), num_starts));
let mut b = Banding::<R>::new(num_slots, seed);
b.add_all(&ordered);
Self { soln: b.solve() }
}
pub fn from_hashable<K: core::hash::Hash>(items: &[K]) -> Self {
let hashes: Vec<u64> = items.iter().map(crate::hash_key).collect();
Self::from_keys_pleated(&hashes)
}
#[inline]
pub fn contains(&self, key: u64) -> bool {
self.soln.contains(key)
}
#[inline]
pub fn contains_hashable<K: core::hash::Hash>(&self, item: &K) -> bool {
self.soln.contains(crate::hash_key(item))
}
pub fn contains_batch(&self, keys: &[u64], out: &mut [bool]) {
self.soln.contains_batch(keys, out);
}
pub fn false_positive_rate(&self) -> f64 {
2f64.powi(-(R as i32))
}
pub fn size_bytes(&self) -> usize {
self.soln.segments().len() * 8
}
pub fn bits_per_key(&self, n: usize) -> f64 {
if n == 0 {
return f64::INFINITY;
}
self.size_bytes() as f64 * 8.0 / n as f64
}
pub fn to_bytes(&self) -> Vec<u8> {
let (num_starts, raw_seed, segs) = self.soln.parts();
let mut buf = crate::format::write_header(
crate::format::FAMILY_HOMOG,
R as u8,
raw_seed,
num_starts,
segs.len() as u64,
);
for &s in segs {
buf.extend_from_slice(&s.to_le_bytes());
}
crate::format::finish(buf)
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, crate::format::DecodeError> {
Self::check_width();
let (hdr, payload) =
crate::format::decode(bytes, crate::format::FAMILY_HOMOG, R as u8, 8, W)?;
let segs: Vec<u64> = payload
.chunks_exact(8)
.map(|c| u64::from_le_bytes(c.try_into().unwrap()))
.collect();
Ok(Self {
soln: Solution::from_parts(hdr.num_starts, hdr.seed, segs),
})
}
#[inline]
fn check_width() {
const { assert!(R >= 1 && R <= 32, "ribbon result width R must be in 1..=32") };
}
#[cfg(all(test, feature = "parallel"))]
pub(crate) fn solution_segments(&self) -> &[u64] {
self.soln.segments()
}
}
#[cfg(feature = "parallel")]
mod parallel;
#[cfg(feature = "parallel")]
impl<const R: usize> Ribbon<R> {
pub fn from_keys_parallel(keys: &[u64], threads: usize) -> Self {
Self::from_keys_parallel_seeded(keys, 0, DEFAULT_WINDOW_SHIFT, threads)
}
pub fn from_keys_parallel_seeded(
keys: &[u64],
seed: u64,
window_shift: u32,
threads: usize,
) -> Self {
Self::check_width();
let _ = window_size(window_shift);
Self {
soln: parallel::from_keys_parallel_seeded::<R>(keys, seed, window_shift, threads),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::banding::solution_fnv;
fn mix64(mut z: u64) -> u64 {
z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
z ^ (z >> 31)
}
fn keys(n: usize, seed: u64) -> Vec<u64> {
let mut s = seed;
(0..n)
.map(|_| {
s = s.wrapping_add(0x9e37_79b9_7f4a_7c15);
mix64(s)
})
.collect()
}
#[test]
fn pleated_build_is_bit_identical_to_arrival() {
for n in [1000usize, 50_000, 250_000] {
let k = keys(n, 0xA11CE);
let arrival = RibbonFilter::from_keys(&k);
let pleated = RibbonFilter::from_keys_pleated(&k);
assert_eq!(
solution_fnv(arrival.soln.segments()),
solution_fnv(pleated.soln.segments()),
"pleated build diverges from arrival at n={n}"
);
}
}
#[test]
fn no_false_negatives_and_plausible_fpr() {
let n = 200_000;
let k = keys(n, 0xA11CE);
let f = RibbonFilter::from_keys_pleated(&k);
assert!(k.iter().all(|&x| f.contains(x)), "false negative");
let absent = keys(200_000, 0xD15EA5E);
let fp = absent
.iter()
.filter(|&&x| f.contains(x ^ 0x5555_5555_5555_5555))
.count();
let fpr = fp as f64 / 200_000.0;
assert!(fpr < 0.02, "FPR {fpr} too high"); assert!(f.bits_per_key(n) < 10.0);
}
}
#[cfg(test)]
mod prod_tests {
use super::*;
#[test]
fn empty_and_tiny_inputs_do_not_panic() {
let f = RibbonFilter::from_keys(&[]);
let _empty_probe = f.contains(12345);
assert!(f.size_bytes() > 0);
let f2 = RibbonFilter::from_keys_pleated(&[7, 42, 1000]);
assert!(f2.contains(7) && f2.contains(42) && f2.contains(1000));
}
#[test]
fn tunable_fpr_scales_with_r() {
use crate::filter::Ribbon;
let k: Vec<u64> = (0..200_000u64)
.map(|i| i.wrapping_mul(0x9e3779b97f4a7c15))
.collect();
let absent: Vec<u64> = (0..200_000u64)
.map(|i| i.wrapping_mul(0x9e3779b97f4a7c15) ^ 0x1)
.collect();
let fpr = |present: &dyn Fn(u64) -> bool| -> f64 {
absent.iter().filter(|&&x| present(x)).count() as f64 / absent.len() as f64
};
let f7 = Ribbon::<7>::from_keys_pleated(&k);
let f10 = Ribbon::<10>::from_keys_pleated(&k);
assert!(k.iter().all(|&x| f7.contains(x)) && k.iter().all(|&x| f10.contains(x)));
let (e7, e10) = (fpr(&|x| f7.contains(x)), fpr(&|x| f10.contains(x)));
assert!(e10 < e7, "r=10 FPR {e10} should be below r=7 {e7}");
assert!(f10.bits_per_key(k.len()) > f7.bits_per_key(k.len()));
}
#[test]
fn roundtrip_serialization_preserves_queries() {
let k: Vec<u64> = (0..100_000u64)
.map(|i| i.wrapping_mul(0x9e3779b97f4a7c15))
.collect();
let f = RibbonFilter::from_keys_pleated(&k);
let bytes = f.to_bytes();
let g = RibbonFilter::from_bytes(&bytes).expect("valid buffer");
assert_eq!(f.size_bytes(), g.size_bytes());
assert!(
k.iter().all(|&x| g.contains(x)),
"false negative after roundtrip"
);
for x in [1u64, 3, 999_999_999, u64::MAX] {
assert_eq!(f.contains(x), g.contains(x));
}
assert!(RibbonFilter::from_bytes(&[0u8; 5]).is_err());
}
}
use crate::banding128::{build_std128, build_std128_pleated, Solution128, W128};
pub struct StdRibbon<const R: usize> {
soln: Solution128<R>,
}
impl<const R: usize> core::fmt::Debug for StdRibbon<R> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("StdRibbon")
.field("w", &128u32)
.field("r", &R)
.field("bytes", &self.size_bytes())
.finish()
}
}
impl<const R: usize> StdRibbon<R> {
pub fn from_keys(keys: &[u64]) -> Result<Self, BuildError> {
Self::check_width();
build_std128::<R>(keys)
.map(|soln| Self { soln })
.ok_or(BuildError::NoSolvingSeed)
}
pub fn from_keys_pleated(keys: &[u64]) -> Result<Self, BuildError> {
Self::check_width();
build_std128_pleated::<R>(keys, DEFAULT_WINDOW_SHIFT)
.map(|soln| Self { soln })
.ok_or(BuildError::NoSolvingSeed)
}
pub fn from_hashable<K: core::hash::Hash>(items: &[K]) -> Result<Self, BuildError> {
let hashes: Vec<u64> = items.iter().map(crate::hash_key).collect();
Self::from_keys_pleated(&hashes)
}
#[cfg(feature = "parallel")]
pub fn from_keys_parallel(keys: &[u64], threads: usize) -> Result<Self, BuildError> {
Self::check_width();
crate::banding128::build_std128_parallel::<R>(keys, DEFAULT_WINDOW_SHIFT, threads)
.map(|soln| Self { soln })
.ok_or(BuildError::NoSolvingSeed)
}
#[inline]
pub fn contains(&self, key: u64) -> bool {
self.soln.contains(key)
}
#[inline]
pub fn contains_hashable<K: core::hash::Hash>(&self, item: &K) -> bool {
self.soln.contains(crate::hash_key(item))
}
pub fn contains_batch(&self, keys: &[u64], out: &mut [bool]) {
self.soln.contains_batch(keys, out);
}
pub fn false_positive_rate(&self) -> f64 {
2f64.powi(-(R as i32))
}
pub fn size_bytes(&self) -> usize {
self.soln.segments().len() * 16
}
pub fn bits_per_key(&self, n: usize) -> f64 {
if n == 0 {
return f64::INFINITY;
}
self.size_bytes() as f64 * 8.0 / n as f64
}
pub fn to_bytes(&self) -> Vec<u8> {
let (num_starts, ordinal_seed, segs) = self.soln.parts();
let mut buf = crate::format::write_header(
crate::format::FAMILY_STD,
R as u8,
ordinal_seed as u64,
num_starts,
segs.len() as u64,
);
for &s in segs {
buf.extend_from_slice(&s.to_le_bytes());
}
crate::format::finish(buf)
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, crate::format::DecodeError> {
Self::check_width();
let (hdr, payload) =
crate::format::decode(bytes, crate::format::FAMILY_STD, R as u8, 16, W128)?;
if hdr.seed >= crate::banding128::SEED_COUNT as u64 {
return Err(crate::format::DecodeError::BadSeed);
}
let segs: Vec<u128> = payload
.chunks_exact(16)
.map(|c| u128::from_le_bytes(c.try_into().unwrap()))
.collect();
Ok(Self {
soln: Solution128::from_parts(hdr.num_starts, hdr.seed as u32, segs),
})
}
#[inline]
fn check_width() {
const { assert!(R >= 1 && R <= 32, "ribbon result width R must be in 1..=32") };
}
}
#[cfg(test)]
mod std128_tests {
use super::*;
use crate::banding128::solution_fnv_128;
fn keys(n: usize) -> Vec<u64> {
let mut s = 0xA11CEu64;
(0..n)
.map(|_| {
s = s.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut z = s;
z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
z ^ (z >> 31)
})
.collect()
}
#[test]
fn std128_pleated_is_bit_identical_to_arrival() {
for n in [5000usize, 100_000, 300_000] {
let k = keys(n);
let a = StdRibbon::<7>::from_keys(&k).unwrap();
let p = StdRibbon::<7>::from_keys_pleated(&k).unwrap();
assert_eq!(
solution_fnv_128(a.soln.segments()),
solution_fnv_128(p.soln.segments()),
"std128 pleated diverges from arrival at n={n}"
);
assert!(
k.iter().all(|&x| p.contains(x)),
"std128 false negative n={n}"
);
}
}
#[cfg(feature = "parallel")]
#[test]
fn std128_parallel_is_bit_identical_to_arrival() {
let k = keys(300_000);
let a = StdRibbon::<7>::from_keys(&k).unwrap();
for t in [2usize, 4, 8] {
let p = StdRibbon::<7>::from_keys_parallel(&k, t).unwrap();
assert_eq!(
solution_fnv_128(a.soln.segments()),
solution_fnv_128(p.soln.segments()),
"std128 parallel (t={t}) diverges"
);
assert!(
k.iter().all(|&x| p.contains(x)),
"std128 parallel false negative t={t}"
);
}
}
#[test]
fn std128_serialization_roundtrip() {
let k = keys(100_000);
let f = StdRibbon::<8>::from_keys_pleated(&k).unwrap();
let g = StdRibbon::<8>::from_bytes(&f.to_bytes()).unwrap();
assert!(k.iter().all(|&x| g.contains(x)));
for x in [1u64, 7, 999, u64::MAX] {
assert_eq!(f.contains(x), g.contains(x));
}
}
}
#[cfg(test)]
mod hashable_tests {
use super::*;
#[test]
fn hashable_string_and_struct_keys() {
let words: Vec<String> = (0..50_000).map(|i| format!("item-{i}")).collect();
let f = RibbonFilter::from_hashable(&words);
assert!(
words.iter().all(|w| f.contains_hashable(w)),
"false negative on strings"
);
let absent = (0..50_000)
.filter(|i| {
let w = format!("absent-{i}");
f.contains_hashable(&w)
})
.count();
assert!((absent as f64 / 50_000.0) < 0.02, "FPR too high on strings");
let pairs: Vec<(u32, u32)> = (0..20_000u32).map(|i| (i, i.wrapping_mul(7))).collect();
let g = StdRibbon::<8>::from_hashable(&pairs).unwrap();
assert!(pairs.iter().all(|p| g.contains_hashable(p)));
}
}
#[cfg(test)]
mod batch_tests {
use super::*;
fn keys(n: usize) -> Vec<u64> {
(0..n as u64)
.map(|i| i.wrapping_mul(0x9e3779b97f4a7c15))
.collect()
}
#[test]
fn batch_query_matches_scalar() {
let k = keys(100_000);
let f = RibbonFilter::from_keys_pleated(&k);
let probes = keys(50_000);
let mut out = vec![false; probes.len()];
f.contains_batch(&probes, &mut out);
assert!(out.iter().zip(&probes).all(|(&o, &p)| o == f.contains(p)));
assert!((f.false_positive_rate() - 2f64.powi(-7)).abs() < 1e-12);
let g = StdRibbon::<8>::from_keys_pleated(&k).unwrap();
let mut out2 = vec![false; probes.len()];
g.contains_batch(&probes, &mut out2);
assert!(out2.iter().zip(&probes).all(|(&o, &p)| o == g.contains(p)));
}
}
#[cfg(test)]
mod soundness_tests {
use super::*;
use crate::format::DecodeError;
fn keys(n: usize) -> Vec<u64> {
(0..n as u64)
.map(|i| i.wrapping_mul(0x9e3779b97f4a7c15))
.collect()
}
#[test]
fn decode_rejects_malformed_and_mismatched_buffers() {
let f = RibbonFilter::from_keys_pleated(&keys(50_000));
let bytes = f.to_bytes();
let g = RibbonFilter::from_bytes(&bytes).unwrap();
assert_eq!(f.size_bytes(), g.size_bytes());
assert_eq!(
RibbonFilter::from_bytes(&[]).unwrap_err(),
DecodeError::TooShort
);
assert_eq!(
RibbonFilter::from_bytes(&bytes[..bytes.len() - 1]).unwrap_err(),
DecodeError::BadChecksum
);
let mut flipped = bytes.clone();
flipped[40] ^= 1;
assert!(RibbonFilter::from_bytes(&flipped).is_err());
assert_eq!(
StdRibbon::<7>::from_bytes(&bytes).unwrap_err(),
DecodeError::WrongFamily
);
assert_eq!(
Ribbon::<8>::from_bytes(&bytes).unwrap_err(),
DecodeError::WrongResultWidth
);
let s = StdRibbon::<7>::from_keys_pleated(&keys(50_000))
.unwrap()
.to_bytes();
assert_eq!(
RibbonFilter::from_bytes(&s).unwrap_err(),
DecodeError::WrongFamily
);
assert!(StdRibbon::<7>::from_bytes(&s)
.unwrap()
.contains(keys(50_000)[0]));
}
#[cfg(feature = "parallel")]
#[test]
fn parallel_handles_adversarial_clustered_starts() {
let mut k: Vec<u64> = Vec::new();
for base in 0..2000u64 {
for j in 0..100u64 {
k.push(base.wrapping_mul(0x1_0000).wrapping_add(j));
}
}
let seq = RibbonFilter::from_keys(&k);
let par = RibbonFilter::from_keys_parallel(&k, 8);
assert!(
k.iter().all(|&x| par.contains(x)),
"adversarial parallel false negative"
);
assert_eq!(
seq.to_bytes(),
par.to_bytes(),
"adversarial parallel not bit-identical"
);
}
}
#[cfg(all(test, miri))]
mod miri_tests {
use super::*;
#[test]
fn decode_and_batch_queries_are_memory_safe() {
let keys: Vec<u64> = (0..64).map(|i| i * 0x1_0001).collect();
let homogeneous = RibbonFilter::from_keys_pleated(&keys);
let decoded = RibbonFilter::from_bytes(&homogeneous.to_bytes()).unwrap();
let mut out = vec![false; keys.len()];
decoded.contains_batch(&keys, &mut out);
assert!(out.into_iter().all(core::convert::identity));
let standard = StdRibbon::<7>::from_keys_pleated(&keys).unwrap();
let decoded = StdRibbon::<7>::from_bytes(&standard.to_bytes()).unwrap();
let mut out = vec![false; keys.len()];
decoded.contains_batch(&keys, &mut out);
assert!(out.into_iter().all(core::convert::identity));
}
}
#[cfg(test)]
mod fpr_stat_tests {
use super::*;
fn keys(n: usize, seed: u64) -> Vec<u64> {
let mut s = seed;
(0..n)
.map(|_| {
s = s.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut z = s;
z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
z ^ (z >> 31)
})
.collect()
}
fn assert_fpr_near<F: Fn(u64) -> bool>(present: F, r: u32) {
let probes = keys(2_000_000, 0xD15EA5E);
let fp = probes
.iter()
.filter(|&&k| present(k ^ 0x5555_5555_5555_5555))
.count();
let measured = fp as f64 / probes.len() as f64;
let expected = 2f64.powi(-(r as i32));
let rel = (measured - expected).abs() / expected;
assert!(
rel < 0.15,
"r={r}: measured FPR {measured:.5} deviates {:.1}% from 2^-r {expected:.5}",
rel * 100.0
);
}
#[test]
fn fpr_matches_two_to_minus_r_statistically() {
let k = keys(500_000, 0xA11CE);
let f7 = Ribbon::<7>::from_keys_pleated(&k);
assert_fpr_near(|x| f7.contains(x), 7);
let f10 = Ribbon::<10>::from_keys_pleated(&k);
assert_fpr_near(|x| f10.contains(x), 10);
let s7 = StdRibbon::<7>::from_keys_pleated(&k).unwrap();
assert_fpr_near(|x| s7.contains(x), 7);
}
}