use std::collections::{BTreeMap, BTreeSet};
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub struct TreeShape {
pub n_time_buckets: u32,
pub n_segments: u32,
pub time_window_seconds: u64,
}
#[derive(Debug, Clone)]
struct Bucket {
segments: Vec<u64>,
directory: BTreeMap<u32, BTreeSet<KeyEntry>>,
}
#[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd)]
pub struct KeyEntry {
pub bucket: Vec<u8>,
pub key: Vec<u8>,
pub vclock: Vec<u8>,
}
impl KeyEntry {
#[must_use]
pub fn hash(&self) -> u64 {
let mut h = FNV1A_OFFSET;
for byte in &self.bucket {
h ^= u64::from(*byte);
h = h.wrapping_mul(FNV1A_PRIME);
}
h ^= 0;
h = h.wrapping_mul(FNV1A_PRIME);
for byte in &self.key {
h ^= u64::from(*byte);
h = h.wrapping_mul(FNV1A_PRIME);
}
h ^= 0;
h = h.wrapping_mul(FNV1A_PRIME);
for byte in &self.vclock {
h ^= u64::from(*byte);
h = h.wrapping_mul(FNV1A_PRIME);
}
h
}
#[must_use]
pub fn segment_id(&self, n_segments: u32) -> u32 {
let h = self.hash();
let mixed = (h >> 32) ^ (h & 0xffff_ffff);
let modded = mixed % u64::from(n_segments);
u32::try_from(modded).expect("invariant: modulo by u32 fits in u32")
}
}
const FNV1A_OFFSET: u64 = 0xcbf2_9ce4_8422_2325;
const FNV1A_PRIME: u64 = 0x0000_0100_0000_01b3;
#[derive(Debug, Clone)]
pub struct Tree {
shape: TreeShape,
buckets: Vec<Bucket>,
}
impl Tree {
#[must_use]
pub fn new(shape: TreeShape) -> Self {
assert!(
shape.n_time_buckets > 0,
"TreeShape::n_time_buckets must be > 0"
);
assert!(shape.n_segments > 0, "TreeShape::n_segments must be > 0");
let buckets = (0..shape.n_time_buckets)
.map(|_| Bucket {
segments: vec![0u64; shape.n_segments as usize],
directory: BTreeMap::new(),
})
.collect();
Self { shape, buckets }
}
#[must_use]
pub fn shape(&self) -> TreeShape {
self.shape
}
#[must_use]
pub fn time_bucket_id(&self, timestamp_seconds: u64) -> u32 {
let window = self.shape.time_window_seconds.max(1);
let bucket = (timestamp_seconds / window) % u64::from(self.shape.n_time_buckets);
u32::try_from(bucket).expect("invariant: modulo by u32 fits in u32")
}
pub fn insert(&mut self, bucket: &[u8], key: &[u8], vclock: &[u8], timestamp_seconds: u64) {
let entry = KeyEntry {
bucket: bucket.to_vec(),
key: key.to_vec(),
vclock: vclock.to_vec(),
};
let tb = self.time_bucket_id(timestamp_seconds);
let seg = entry.segment_id(self.shape.n_segments);
let row = &mut self.buckets[tb as usize];
let dir_set = row.directory.entry(seg).or_default();
if dir_set.insert(entry.clone()) {
row.segments[seg as usize] ^= entry.hash();
}
}
pub fn remove(&mut self, bucket: &[u8], key: &[u8], vclock: &[u8], timestamp_seconds: u64) {
let entry = KeyEntry {
bucket: bucket.to_vec(),
key: key.to_vec(),
vclock: vclock.to_vec(),
};
let tb = self.time_bucket_id(timestamp_seconds);
let seg = entry.segment_id(self.shape.n_segments);
let row = &mut self.buckets[tb as usize];
if let Some(set) = row.directory.get_mut(&seg) {
if set.remove(&entry) {
row.segments[seg as usize] ^= entry.hash();
if set.is_empty() {
row.directory.remove(&seg);
}
}
}
}
pub fn update(
&mut self,
bucket: &[u8],
key: &[u8],
old_vclock: &[u8],
new_vclock: &[u8],
old_timestamp: u64,
new_timestamp: u64,
) {
self.remove(bucket, key, old_vclock, old_timestamp);
self.insert(bucket, key, new_vclock, new_timestamp);
}
#[must_use]
pub fn roots(&self) -> Vec<(u32, u64)> {
self.buckets
.iter()
.enumerate()
.map(|(i, b)| {
let root = b.segments.iter().copied().fold(0u64, |a, x| a ^ x);
let i = u32::try_from(i)
.expect("invariant: time bucket index fits in u32 by construction");
(i, root)
})
.collect()
}
pub fn segments(&self, time_bucket: u32) -> Result<Vec<(u32, u64)>, TreeError> {
let row = self
.buckets
.get(time_bucket as usize)
.ok_or(TreeError::TimeBucketOutOfRange(time_bucket))?;
Ok(row
.segments
.iter()
.copied()
.enumerate()
.map(|(i, h)| {
let i =
u32::try_from(i).expect("invariant: segment index fits in u32 by construction");
(i, h)
})
.collect())
}
pub fn keys_in_segment(
&self,
time_bucket: u32,
segment: u32,
) -> Result<Vec<KeyEntry>, TreeError> {
let row = self
.buckets
.get(time_bucket as usize)
.ok_or(TreeError::TimeBucketOutOfRange(time_bucket))?;
Ok(row
.directory
.get(&segment)
.map(|set| set.iter().cloned().collect())
.unwrap_or_default())
}
#[must_use]
pub fn diverging_time_buckets(local: &[(u32, u64)], remote: &[(u32, u64)]) -> Vec<u32> {
let remote_map: BTreeMap<u32, u64> = remote.iter().copied().collect();
let mut out = Vec::new();
for (idx, local_root) in local {
if let Some(remote_root) = remote_map.get(idx) {
if remote_root != local_root {
out.push(*idx);
}
} else {
out.push(*idx);
}
}
out
}
#[must_use]
pub fn diverging_segments(local: &[(u32, u64)], remote: &[(u32, u64)]) -> Vec<u32> {
let remote_map: BTreeMap<u32, u64> = remote.iter().copied().collect();
let mut out = Vec::new();
for (idx, local_hash) in local {
if let Some(remote_hash) = remote_map.get(idx) {
if remote_hash != local_hash {
out.push(*idx);
}
} else {
out.push(*idx);
}
}
out
}
}
#[derive(Debug, thiserror::Error)]
pub enum TreeError {
#[error("time bucket {0} out of range")]
TimeBucketOutOfRange(u32),
#[error("segment {0} out of range")]
SegmentOutOfRange(u32),
}
impl Tree {
pub(crate) fn install_segment(
&mut self,
time_bucket: u32,
segment: u32,
hash: u64,
entries: Vec<KeyEntry>,
) -> Result<(), TreeError> {
let row = self
.buckets
.get_mut(time_bucket as usize)
.ok_or(TreeError::TimeBucketOutOfRange(time_bucket))?;
if segment >= self.shape.n_segments {
return Err(TreeError::SegmentOutOfRange(segment));
}
row.segments[segment as usize] = hash;
let set: BTreeSet<KeyEntry> = entries.into_iter().collect();
if set.is_empty() {
row.directory.remove(&segment);
} else {
row.directory.insert(segment, set);
}
Ok(())
}
pub(crate) fn collect_nonempty_segments(&self) -> Vec<(u32, u32, u64, Vec<KeyEntry>)> {
let mut out = Vec::new();
for (tb_idx, row) in self.buckets.iter().enumerate() {
let tb = u32::try_from(tb_idx)
.expect("invariant: time bucket index fits in u32 by construction");
for (seg, set) in &row.directory {
if set.is_empty() {
continue;
}
let hash = row.segments[*seg as usize];
let entries: Vec<KeyEntry> = set.iter().cloned().collect();
out.push((tb, *seg, hash, entries));
}
}
out
}
}
#[cfg(test)]
mod tests {
use super::*;
fn shape() -> TreeShape {
TreeShape {
n_time_buckets: 4,
n_segments: 64,
time_window_seconds: 60,
}
}
#[test]
fn empty_tree_roots_are_zero() {
let t = Tree::new(shape());
for (_, root) in t.roots() {
assert_eq!(root, 0);
}
}
#[test]
fn merkle_round_trip_localizes_one_leaf() {
let mut a = Tree::new(shape());
let mut b = Tree::new(shape());
for i in 0..1000u32 {
let key = format!("k{i}");
let vc = format!("vc{i}");
a.insert(b"users", key.as_bytes(), vc.as_bytes(), 0);
b.insert(b"users", key.as_bytes(), vc.as_bytes(), 0);
}
assert_eq!(a.roots(), b.roots());
b.update(b"users", b"k42", b"vc42", b"vc42-updated", 0, 0);
let dr = Tree::diverging_time_buckets(&a.roots(), &b.roots());
assert_eq!(dr.len(), 1, "only one time bucket should diverge");
let tb = dr[0];
let ds = Tree::diverging_segments(&a.segments(tb).unwrap(), &b.segments(tb).unwrap());
assert!(
(1..=2).contains(&ds.len()),
"expected 1 or 2 diverging segments, got {ds:?}"
);
let mut found_local_old = false;
let mut found_remote_new = false;
for seg in &ds {
for entry in a.keys_in_segment(tb, *seg).unwrap() {
if entry.key == b"k42" && entry.vclock == b"vc42" {
found_local_old = true;
}
}
for entry in b.keys_in_segment(tb, *seg).unwrap() {
if entry.key == b"k42" && entry.vclock == b"vc42-updated" {
found_remote_new = true;
}
}
}
assert!(found_local_old);
assert!(found_remote_new);
}
#[test]
fn xor_is_its_own_inverse() {
let mut t = Tree::new(shape());
let baseline = t.roots();
t.insert(b"b", b"k", b"vc", 0);
assert_ne!(t.roots(), baseline);
t.remove(b"b", b"k", b"vc", 0);
assert_eq!(t.roots(), baseline);
}
#[test]
fn duplicate_insert_is_idempotent() {
let mut t = Tree::new(shape());
t.insert(b"b", b"k", b"vc", 0);
let after_one = t.roots();
t.insert(b"b", b"k", b"vc", 0);
assert_eq!(t.roots(), after_one);
}
#[test]
fn time_bucket_id_rolls_over() {
let t = Tree::new(shape());
let n = t.shape.n_time_buckets;
let w = t.shape.time_window_seconds;
assert_eq!(t.time_bucket_id(0), 0);
assert_eq!(t.time_bucket_id(w), 1);
assert_eq!(t.time_bucket_id(u64::from(n) * w), 0);
}
#[test]
fn segments_out_of_range_errors() {
let t = Tree::new(shape());
assert!(t.segments(999).is_err());
}
#[test]
fn diverging_buckets_handles_size_mismatch() {
let local = vec![(0u32, 1u64), (1, 2), (2, 3)];
let remote = vec![(0u32, 1u64), (1, 99)];
let d = Tree::diverging_time_buckets(&local, &remote);
assert_eq!(d, vec![1, 2]);
}
}