use crate::hash::{coeff_row, ribbon_hash, start};
pub const W: usize = 64;
pub struct Banding<const R: usize> {
coeff_rows: Vec<u64>,
num_starts: u64,
raw_seed: u64,
}
impl<const R: usize> Banding<R> {
pub fn new(num_slots: usize, raw_seed: u64) -> Self {
assert_eq!(num_slots % W, 0, "num_slots must be a multiple of 64");
Self {
coeff_rows: vec![0u64; num_slots],
num_starts: (num_slots - W + 1) as u64,
raw_seed,
}
}
#[cfg(feature = "parallel")]
pub fn num_slots(&self) -> usize {
self.coeff_rows.len()
}
#[cfg(feature = "parallel")]
pub fn coeff_rows_mut(&mut self) -> &mut [u64] {
&mut self.coeff_rows
}
pub fn add(&mut self, key: u64) -> bool {
let h = ribbon_hash(key, self.raw_seed);
let mut i = start(h, self.num_starts) as usize;
let mut cr = coeff_row(h);
loop {
let cr_at_i = self.coeff_rows[i];
if cr_at_i == 0 {
self.coeff_rows[i] = cr;
return true;
}
cr ^= cr_at_i;
if cr == 0 {
return true;
}
let tz = cr.trailing_zeros() as usize;
i += tz;
cr >>= tz;
}
}
pub fn add_all(&mut self, keys: &[u64]) {
for &k in keys {
self.add(k);
}
}
pub fn solve(&self) -> Solution<R> {
Solution {
segments: self.back_substitute(),
num_starts: self.num_starts,
raw_seed: self.raw_seed,
}
}
pub fn back_substitute(&self) -> Vec<u64> {
let num_slots = self.coeff_rows.len();
let num_blocks = num_slots / W;
let num_segments = num_blocks * R;
let mut data = vec![0u64; num_segments];
let mut state = [0u64; R];
for block in (0..num_blocks).rev() {
let start_slot = block * W;
for i in (start_slot..start_slot + W).rev() {
let cr = self.coeff_rows[i];
let rr: u32 = if cr == 0 {
(i as u64).wrapping_mul(0x9E37_79B1_85EB_CA87) as u32
} else {
0
};
for (j, sj) in state.iter_mut().enumerate() {
let tmp = *sj << 1;
let bit = (((tmp & cr).count_ones() & 1) as u64) ^ (((rr >> j) & 1) as u64);
*sj = tmp | bit;
}
}
let segment_num = block * R;
data[segment_num..segment_num + R].copy_from_slice(&state);
}
data
}
}
pub struct Solution<const R: usize> {
segments: Vec<u64>,
num_starts: u64,
raw_seed: u64,
}
impl<const R: usize> Solution<R> {
pub fn segments(&self) -> &[u64] {
&self.segments
}
pub(crate) fn parts(&self) -> (u64, u64, &[u64]) {
(self.num_starts, self.raw_seed, &self.segments)
}
pub(crate) fn from_parts(num_starts: u64, raw_seed: u64, segments: Vec<u64>) -> Self {
Self {
segments,
num_starts,
raw_seed,
}
}
pub fn contains_batch(&self, keys: &[u64], out: &mut [bool]) {
assert_eq!(keys.len(), out.len());
const STRIDE: usize = 32;
for (kc, oc) in keys.chunks(STRIDE).zip(out.chunks_mut(STRIDE)) {
for &k in kc {
let s = start(ribbon_hash(k, self.raw_seed), self.num_starts) as usize;
let seg = (s / W) * R;
#[cfg(all(target_arch = "x86_64", not(miri)))]
if let Some(p) = self.segments.get(seg) {
unsafe {
core::arch::x86_64::_mm_prefetch(
p as *const _ as *const i8,
core::arch::x86_64::_MM_HINT_T0,
);
}
}
#[cfg(any(not(target_arch = "x86_64"), miri))]
let _ = seg;
}
for (k, o) in kc.iter().zip(oc.iter_mut()) {
*o = self.contains(*k);
}
}
}
pub fn contains(&self, key: u64) -> bool {
let h = ribbon_hash(key, self.raw_seed);
let start_slot = start(h, self.num_starts) as usize;
let start_block = start_slot / W;
let segment_num = start_block * R; let start_bit = start_slot % W;
let cr = coeff_row(h);
let cr_left = cr << start_bit;
let cr_right = cr >> ((W - start_bit) % W);
let maybe = if start_bit != 0 { R } else { 0 };
for i in 0..R {
let soln = (self.segments[segment_num + i] & cr_left)
| (self.segments[segment_num + maybe + i] & cr_right);
if (soln.count_ones() & 1) != 0 {
return false;
}
}
true
}
}
#[cfg(test)]
pub fn solution_fnv(segments: &[u64]) -> u64 {
let mut fnv = 0xcbf2_9ce4_8422_2325u64;
for &seg in segments {
for b in seg.to_le_bytes() {
fnv = (fnv ^ b as u64).wrapping_mul(0x0000_0100_0000_01b3);
}
}
fnv
}
#[cfg(test)]
mod tests {
use super::*;
use std::path::Path;
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()
}
fn vec_field(value: &serde_json::Value, name: &str) -> u64 {
if name == "soln_fnv" {
let hex = value[name]
.as_str()
.expect("soln_fnv must be a hexadecimal string");
u64::from_str_radix(hex, 16).expect("soln_fnv must be valid hexadecimal")
} else {
value[name]
.as_u64()
.unwrap_or_else(|| panic!("missing or non-u64 reference field {name}"))
}
}
fn load() -> serde_json::Value {
let text = std::fs::read_to_string(
Path::new(env!("CARGO_MANIFEST_DIR")).join("tests/vectors/homog_w64_r7.json"),
)
.expect("reference vectors must be present");
serde_json::from_str(&text).expect("reference vectors must be valid JSON")
}
fn by_r(json: &serde_json::Value, r: usize) -> &serde_json::Value {
&json["by_r"][r.to_string()]
}
fn gate_build<const R: usize>(j: &serde_json::Value, r: usize) {
let obj = by_r(j, r);
let n = vec_field(j, "build_n") as usize;
let num_slots = vec_field(obj, "num_slots") as usize;
let expect = vec_field(obj, "soln_fnv");
let mut b = Banding::<R>::new(num_slots, 0);
b.add_all(&keys(n, 0xA11CE));
assert_eq!(
solution_fnv(&b.back_substitute()),
expect,
"solution fingerprint diverges from reference at r={r}"
);
}
#[test]
fn build_fingerprint_matches_reference_all_r() {
let j = load();
gate_build::<5>(&j, 5);
gate_build::<7>(&j, 7);
gate_build::<8>(&j, 8);
gate_build::<10>(&j, 10);
}
fn vec_array(json: &serde_json::Value, name: &str) -> Vec<u8> {
json[name]
.as_array()
.unwrap_or_else(|| panic!("reference field {name} must be an array"))
.iter()
.map(|value| {
u8::try_from(
value
.as_u64()
.unwrap_or_else(|| panic!("reference field {name} must contain integers")),
)
.unwrap_or_else(|_| panic!("reference field {name} contains a non-u8 value"))
})
.collect()
}
fn gate_query<const R: usize>(j: &serde_json::Value, r: usize) {
let obj = by_r(j, r);
let n = vec_field(j, "build_n") as usize;
let num_slots = vec_field(obj, "num_slots") as usize;
let bk = keys(n, 0xA11CE);
let mut b = Banding::<R>::new(num_slots, 0);
b.add_all(&bk);
let soln = b.solve();
for (i, &want) in vec_array(obj, "present").iter().enumerate() {
assert_eq!(
soln.contains(bk[i * 37 % n]) as u8,
want,
"present r={r} i={i}"
);
}
let ak = keys(200, 0xD15EA5E);
for (i, &want) in vec_array(obj, "absent").iter().enumerate() {
assert_eq!(
soln.contains(ak[i] ^ 0x5555_5555_5555_5555) as u8,
want,
"absent r={r} i={i}"
);
}
assert!(bk.iter().all(|&k| soln.contains(k)), "false negative r={r}");
}
#[test]
fn query_outcomes_match_reference_all_r() {
let j = load();
gate_query::<5>(&j, 5);
gate_query::<7>(&j, 7);
gate_query::<8>(&j, 8);
gate_query::<10>(&j, 10);
}
#[test]
fn no_reduction_leaves_low_bit_set() {
let mut b = Banding::<7>::new(64 * 32, 0);
assert!(b.add(12345));
assert!(b
.coeff_rows
.iter()
.filter(|&&c| c != 0)
.all(|&c| c & 1 == 1));
}
}