use std::collections::{BTreeMap, BTreeSet};
use crate::datatypes::set::Tag;
use crate::datatypes::{ActorId, Crdt};
#[derive(Clone, Debug, Default, Eq, PartialEq)]
struct ElementState {
adds: BTreeSet<Tag>,
removes: BTreeSet<Tag>,
}
impl ElementState {
fn is_present(&self) -> bool {
self.adds.iter().any(|t| !self.removes.contains(t))
}
fn join(&mut self, other: &Self) {
self.adds.extend(other.adds.iter().cloned());
self.removes.extend(other.removes.iter().cloned());
}
fn is_empty(&self) -> bool {
self.adds.is_empty() && self.removes.is_empty()
}
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct OrSetDelta {
fragment: BTreeMap<Vec<u8>, ElementState>,
actor_counters: BTreeMap<ActorId, u64>,
}
impl OrSetDelta {
#[must_use]
pub fn is_empty(&self) -> bool {
self.fragment.values().all(ElementState::is_empty) && self.actor_counters.is_empty()
}
pub fn join(&mut self, other: &Self) {
for (actor, &count) in &other.actor_counters {
let entry = self.actor_counters.entry(actor.clone()).or_insert(0);
*entry = (*entry).max(count);
}
for (element, state) in &other.fragment {
self.fragment
.entry(element.clone())
.or_default()
.join(state);
}
}
#[must_use]
pub fn wire_len(&self) -> usize {
let mut n = 0;
for actor in self.actor_counters.keys() {
n += actor.dc.len() + actor.peer.len() + 8;
}
for (element, state) in &self.fragment {
n += element.len();
n += (state.adds.len() + state.removes.len()) * tag_wire_len();
}
n
}
}
fn tag_wire_len() -> usize {
8 + 16
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct DeltaOrSet {
elements: BTreeMap<Vec<u8>, ElementState>,
actor_counters: BTreeMap<ActorId, u64>,
}
impl DeltaOrSet {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn add(&mut self, actor: &ActorId, element: impl Into<Vec<u8>>) -> OrSetDelta {
let element = element.into();
let counter = self.actor_counters.entry(actor.clone()).or_insert(0);
*counter = counter
.checked_add(1)
.expect("delta-or-set counter overflow");
let tag = Tag {
actor: actor.clone(),
counter: *counter,
};
let entry = self.elements.entry(element.clone()).or_default();
entry.adds.insert(tag.clone());
let mut frag = ElementState::default();
frag.adds.insert(tag);
let mut delta = OrSetDelta::default();
delta.fragment.insert(element, frag);
delta.actor_counters.insert(actor.clone(), *counter);
delta
}
pub fn remove(&mut self, element: &[u8]) -> OrSetDelta {
let mut delta = OrSetDelta::default();
if let Some(state) = self.elements.get_mut(element) {
let observed: Vec<Tag> = state.adds.iter().cloned().collect();
if observed.is_empty() {
return delta;
}
let mut frag = ElementState::default();
for tag in observed {
state.removes.insert(tag.clone());
frag.removes.insert(tag);
}
delta.fragment.insert(element.to_vec(), frag);
}
delta
}
pub fn merge_delta(&mut self, delta: &OrSetDelta) {
for (actor, &count) in &delta.actor_counters {
let entry = self.actor_counters.entry(actor.clone()).or_insert(0);
*entry = (*entry).max(count);
}
for (element, frag) in &delta.fragment {
self.elements.entry(element.clone()).or_default().join(frag);
}
}
#[must_use]
pub fn contains(&self, element: &[u8]) -> bool {
self.elements
.get(element)
.is_some_and(ElementState::is_present)
}
#[must_use]
pub fn wire_len(&self) -> usize {
let mut n = 0;
for actor in self.actor_counters.keys() {
n += actor.dc.len() + actor.peer.len() + 8;
}
for (element, state) in &self.elements {
n += element.len();
n += (state.adds.len() + state.removes.len()) * tag_wire_len();
}
n
}
}
impl Crdt for DeltaOrSet {
type Value = BTreeSet<Vec<u8>>;
fn merge(&mut self, other: &Self) {
for (actor, &count) in &other.actor_counters {
let entry = self.actor_counters.entry(actor.clone()).or_insert(0);
*entry = (*entry).max(count);
}
for (element, state) in &other.elements {
self.elements
.entry(element.clone())
.or_default()
.join(state);
}
}
fn value(&self) -> BTreeSet<Vec<u8>> {
self.elements
.iter()
.filter(|(_e, s)| s.is_present())
.map(|(e, _s)| e.clone())
.collect()
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct BufferedDelta {
pub seq: u64,
pub delta: OrSetDelta,
}
#[derive(Clone, Debug, Default)]
pub struct DeltaBuffer {
next_seq: u64,
deltas: Vec<BufferedDelta>,
acked: BTreeMap<String, u64>,
}
impl DeltaBuffer {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn record(&mut self, delta: OrSetDelta) -> Option<u64> {
if delta.is_empty() {
return None;
}
let seq = self.next_seq;
self.next_seq += 1;
self.deltas.push(BufferedDelta { seq, delta });
Some(seq)
}
#[must_use]
pub fn high_water(&self) -> Option<u64> {
self.next_seq.checked_sub(1)
}
#[must_use]
pub fn interval_since(&self, since: u64) -> Option<OrSetDelta> {
let mut acc = OrSetDelta::default();
let mut any = false;
for buffered in &self.deltas {
if buffered.seq >= since {
acc.join(&buffered.delta);
any = true;
}
}
any.then_some(acc)
}
#[must_use]
pub fn knows_peer(&self, peer: &str) -> bool {
self.acked.contains_key(peer)
}
#[must_use]
pub fn next_needed(&self, peer: &str) -> Option<u64> {
self.acked.get(peer).map(|acked| acked + 1)
}
pub fn ack(&mut self, peer: &str, seq: u64) {
let entry = self.acked.entry(peer.to_string()).or_insert(0);
*entry = (*entry).max(seq);
}
pub fn compact(&mut self) {
let Some(min_acked) = self.acked.values().min().copied() else {
return;
};
self.deltas.retain(|b| b.seq > min_acked);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn aid(name: &str) -> ActorId {
ActorId::new("dc1", name)
}
#[test]
fn add_then_contains() {
let a = aid("a");
let mut s = DeltaOrSet::new();
s.add(&a, b"x".to_vec());
assert!(s.contains(b"x"));
}
#[test]
fn delta_reproduces_add_on_fresh_replica() {
let a = aid("a");
let mut src = DeltaOrSet::new();
let d = src.add(&a, b"x".to_vec());
let mut dst = DeltaOrSet::new();
dst.merge_delta(&d);
assert!(dst.contains(b"x"));
assert_eq!(src.value(), dst.value());
}
#[test]
fn delta_reproduces_remove_on_fresh_replica() {
let a = aid("a");
let mut src = DeltaOrSet::new();
let add = src.add(&a, b"x".to_vec());
let rem = src.remove(b"x");
let mut dst = DeltaOrSet::new();
dst.merge_delta(&add);
dst.merge_delta(&rem);
assert!(!dst.contains(b"x"));
assert_eq!(src.value(), dst.value());
}
#[test]
fn out_of_order_delivery_still_converges() {
let a = aid("a");
let mut src = DeltaOrSet::new();
let add = src.add(&a, b"x".to_vec());
let rem = src.remove(b"x");
let mut dst = DeltaOrSet::new();
dst.merge_delta(&rem);
dst.merge_delta(&add);
assert_eq!(src.value(), dst.value());
assert!(!dst.contains(b"x"));
}
#[test]
fn duplicate_delta_is_idempotent() {
let a = aid("a");
let mut src = DeltaOrSet::new();
let d = src.add(&a, b"x".to_vec());
let mut dst = DeltaOrSet::new();
dst.merge_delta(&d);
let once = dst.value();
dst.merge_delta(&d);
dst.merge_delta(&d);
assert_eq!(dst.value(), once);
}
#[test]
fn concurrent_remove_loses_to_concurrent_add() {
let a = aid("a");
let b = aid("b");
let mut shared = DeltaOrSet::new();
let seed = shared.add(&a, b"x".to_vec());
let mut left = DeltaOrSet::new();
left.merge_delta(&seed);
let left_rem = left.remove(b"x");
let mut right = DeltaOrSet::new();
right.merge_delta(&seed);
let right_add = right.add(&b, b"x".to_vec());
left.merge_delta(&right_add);
right.merge_delta(&left_rem);
assert!(left.contains(b"x"));
assert!(right.contains(b"x"));
assert_eq!(left.value(), right.value());
}
#[test]
fn interval_joins_buffered_deltas() {
let a = aid("a");
let mut src = DeltaOrSet::new();
let mut buf = DeltaBuffer::new();
buf.record(src.add(&a, b"x".to_vec()));
buf.record(src.add(&a, b"y".to_vec()));
buf.record(src.add(&a, b"z".to_vec()));
let interval = buf.interval_since(1).expect("interval");
let mut dst = DeltaOrSet::new();
dst.merge_delta(&interval);
assert!(!dst.contains(b"x"));
assert!(dst.contains(b"y"));
assert!(dst.contains(b"z"));
}
#[test]
fn empty_delta_not_buffered() {
let mut buf = DeltaBuffer::new();
let mut s = DeltaOrSet::new();
let empty = s.remove(b"absent");
assert!(empty.is_empty());
assert_eq!(buf.record(empty), None);
assert_eq!(buf.high_water(), None);
}
#[test]
fn compact_drops_fully_acked_deltas() {
let a = aid("a");
let mut src = DeltaOrSet::new();
let mut buf = DeltaBuffer::new();
buf.record(src.add(&a, b"x".to_vec()));
buf.record(src.add(&a, b"y".to_vec()));
buf.ack("peer-b", 0);
buf.compact();
assert!(buf.interval_since(0).is_some());
let interval = buf.interval_since(0).unwrap();
let mut dst = DeltaOrSet::new();
dst.merge_delta(&interval);
assert!(!dst.contains(b"x"));
assert!(dst.contains(b"y"));
}
#[test]
fn delta_merge_matches_full_state_merge() {
let a = aid("a");
let b = aid("b");
let mut producer = DeltaOrSet::new();
let d1 = producer.add(&a, b"x".to_vec());
let d2 = producer.add(&b, b"y".to_vec());
let d3 = producer.remove(b"x");
let mut via_deltas = DeltaOrSet::new();
via_deltas.merge_delta(&d1);
via_deltas.merge_delta(&d2);
via_deltas.merge_delta(&d3);
let mut via_state = DeltaOrSet::new();
via_state.merge(&producer);
assert_eq!(via_deltas.value(), via_state.value());
assert_eq!(via_deltas, via_state);
}
#[test]
fn first_contact_has_no_ack_point() {
let buf = DeltaBuffer::new();
assert!(!buf.knows_peer("peer-b"));
assert_eq!(buf.next_needed("peer-b"), None);
}
}