use std::collections::{BTreeMap, BTreeSet};
use crate::datatypes::set::Tag;
use crate::datatypes::{ActorId, Crdt, EwFlag, LwwRegister, OrSet, PnCounter};
#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub enum FieldType {
Counter = 1,
OrSet = 2,
LwwRegister = 3,
EwFlag = 4,
NestedMap = 5,
}
impl FieldType {
#[must_use]
pub fn from_wire(code: i32) -> Option<Self> {
match code {
1 => Some(FieldType::Counter),
2 => Some(FieldType::OrSet),
3 => Some(FieldType::LwwRegister),
4 => Some(FieldType::EwFlag),
5 => Some(FieldType::NestedMap),
_ => None,
}
}
#[must_use]
pub fn to_wire(self) -> i32 {
self as i32
}
}
#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct FieldKey {
pub name: Vec<u8>,
pub field_type: FieldType,
}
impl FieldKey {
pub fn new(name: impl Into<Vec<u8>>, field_type: FieldType) -> Self {
Self {
name: name.into(),
field_type,
}
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum FieldValue {
Counter(PnCounter),
OrSet(OrSet),
LwwRegister(LwwRegister),
EwFlag(EwFlag),
NestedMap(Box<Map>),
}
impl FieldValue {
fn empty_for(field_type: FieldType) -> Self {
match field_type {
FieldType::Counter => FieldValue::Counter(PnCounter::new()),
FieldType::OrSet => FieldValue::OrSet(OrSet::new()),
FieldType::LwwRegister => FieldValue::LwwRegister(LwwRegister::new()),
FieldType::EwFlag => FieldValue::EwFlag(EwFlag::new()),
FieldType::NestedMap => FieldValue::NestedMap(Box::new(Map::new())),
}
}
fn merge(&mut self, other: &FieldValue) {
match (self, other) {
(FieldValue::Counter(a), FieldValue::Counter(b)) => a.merge(b),
(FieldValue::OrSet(a), FieldValue::OrSet(b)) => a.merge(b),
(FieldValue::LwwRegister(a), FieldValue::LwwRegister(b)) => a.merge(b),
(FieldValue::EwFlag(a), FieldValue::EwFlag(b)) => a.merge(b),
(FieldValue::NestedMap(a), FieldValue::NestedMap(b)) => a.merge(b),
_ => {}
}
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum NestedOp {
Counter(i64),
SetAdd(Vec<u8>),
SetRemove(Vec<u8>),
RegisterAssign {
value: Vec<u8>,
ts_micros: u64,
},
Flag(bool),
Map(Box<MapOp>),
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum MapOp {
Update {
field: FieldKey,
op: NestedOp,
},
Remove {
field: FieldKey,
},
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
struct FieldEntry {
adds: BTreeSet<Tag>,
removes: BTreeSet<Tag>,
value: FieldValue,
}
impl FieldEntry {
fn new(field_type: FieldType) -> Self {
Self {
adds: BTreeSet::new(),
removes: BTreeSet::new(),
value: FieldValue::empty_for(field_type),
}
}
fn is_present(&self) -> bool {
self.adds.iter().any(|t| !self.removes.contains(t))
}
fn merge(&mut self, other: &FieldEntry) {
for tag in &other.adds {
self.adds.insert(tag.clone());
}
for tag in &other.removes {
self.removes.insert(tag.clone());
}
self.value.merge(&other.value);
}
}
impl Default for FieldValue {
fn default() -> Self {
FieldValue::Counter(PnCounter::new())
}
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct Map {
fields: BTreeMap<FieldKey, FieldEntry>,
actor_counters: BTreeMap<ActorId, u64>,
}
impl Map {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn apply(&mut self, actor: &ActorId, op: &MapOp) {
match op {
MapOp::Update { field, op } => self.apply_update(actor, field, op),
MapOp::Remove { field } => self.apply_remove(field),
}
}
fn apply_update(&mut self, actor: &ActorId, field: &FieldKey, op: &NestedOp) {
let counter = self.actor_counters.entry(actor.clone()).or_insert(0);
*counter = counter.checked_add(1).expect("map counter overflow");
let tag = Tag {
actor: actor.clone(),
counter: *counter,
};
let entry = self
.fields
.entry(field.clone())
.or_insert_with(|| FieldEntry::new(field.field_type));
entry.adds.insert(tag);
Self::apply_nested(actor, op, &mut entry.value);
}
fn apply_nested(actor: &ActorId, op: &NestedOp, value: &mut FieldValue) {
match (op, value) {
(NestedOp::Counter(delta), FieldValue::Counter(c)) => c.apply(actor, *delta),
(NestedOp::SetAdd(elt), FieldValue::OrSet(s)) => {
s.add(actor, elt.clone());
}
(NestedOp::SetRemove(elt), FieldValue::OrSet(s)) => s.remove(elt),
(
NestedOp::RegisterAssign {
value: v,
ts_micros,
},
FieldValue::LwwRegister(r),
) => {
r.assign(actor, *ts_micros, v.clone());
}
(NestedOp::Flag(true), FieldValue::EwFlag(f)) => {
f.enable(actor);
}
(NestedOp::Flag(false), FieldValue::EwFlag(f)) => f.disable(),
(NestedOp::Map(inner), FieldValue::NestedMap(m)) => m.apply(actor, inner),
_ => {
}
}
}
fn apply_remove(&mut self, field: &FieldKey) {
if let Some(entry) = self.fields.get_mut(field) {
for tag in entry.adds.clone() {
entry.removes.insert(tag);
}
}
}
#[must_use]
pub fn contains(&self, field: &FieldKey) -> bool {
self.fields.get(field).is_some_and(FieldEntry::is_present)
}
#[must_use]
pub fn get(&self, field: &FieldKey) -> Option<&FieldValue> {
self.fields
.get(field)
.filter(|e| e.is_present())
.map(|e| &e.value)
}
pub fn iter(&self) -> impl Iterator<Item = (&FieldKey, &FieldValue)> {
self.fields
.iter()
.filter(|(_, e)| e.is_present())
.map(|(k, e)| (k, &e.value))
}
}
impl Crdt for Map {
type Value = BTreeMap<FieldKey, FieldValue>;
fn merge(&mut self, other: &Self) {
for (actor, &count) in &other.actor_counters {
let entry = self.actor_counters.entry(actor.clone()).or_insert(0);
if *entry < count {
*entry = count;
}
}
for (key, other_entry) in &other.fields {
let entry = self
.fields
.entry(key.clone())
.or_insert_with(|| FieldEntry::new(key.field_type));
entry.merge(other_entry);
}
}
fn value(&self) -> BTreeMap<FieldKey, FieldValue> {
self.fields
.iter()
.filter_map(|(k, e)| {
if e.is_present() {
Some((k.clone(), e.value.clone()))
} else {
None
}
})
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn aid(name: &str) -> ActorId {
ActorId::new("dc1", name)
}
fn counter_field(name: &str) -> FieldKey {
FieldKey::new(name.as_bytes(), FieldType::Counter)
}
fn set_field(name: &str) -> FieldKey {
FieldKey::new(name.as_bytes(), FieldType::OrSet)
}
fn flag_field(name: &str) -> FieldKey {
FieldKey::new(name.as_bytes(), FieldType::EwFlag)
}
fn register_field(name: &str) -> FieldKey {
FieldKey::new(name.as_bytes(), FieldType::LwwRegister)
}
fn map_field(name: &str) -> FieldKey {
FieldKey::new(name.as_bytes(), FieldType::NestedMap)
}
#[test]
fn fresh_map_is_empty() {
let m = Map::new();
assert!(m.value().is_empty());
}
#[test]
fn update_creates_field_with_correct_type() {
let a = aid("a");
let mut m = Map::new();
let f = counter_field("hits");
m.apply(
&a,
&MapOp::Update {
field: f.clone(),
op: NestedOp::Counter(5),
},
);
assert!(m.contains(&f));
match m.get(&f) {
Some(FieldValue::Counter(c)) => assert_eq!(c.value(), 5),
_ => panic!("counter field missing or wrong type"),
}
}
#[test]
fn remove_drops_field_from_value() {
let a = aid("a");
let mut m = Map::new();
let f = flag_field("on");
m.apply(
&a,
&MapOp::Update {
field: f.clone(),
op: NestedOp::Flag(true),
},
);
assert!(m.contains(&f));
m.apply(&a, &MapOp::Remove { field: f.clone() });
assert!(!m.contains(&f));
}
#[test]
fn three_typed_fields_round_trip_through_value() {
let a = aid("a");
let mut m = Map::new();
let cf = counter_field("c");
let sf = set_field("s");
let rf = register_field("r");
m.apply(
&a,
&MapOp::Update {
field: cf.clone(),
op: NestedOp::Counter(7),
},
);
m.apply(
&a,
&MapOp::Update {
field: sf.clone(),
op: NestedOp::SetAdd(b"x".to_vec()),
},
);
m.apply(
&a,
&MapOp::Update {
field: rf.clone(),
op: NestedOp::RegisterAssign {
value: b"hello".to_vec(),
ts_micros: 100,
},
},
);
let v = m.value();
assert_eq!(v.len(), 3);
match v.get(&cf) {
Some(FieldValue::Counter(c)) => assert_eq!(c.value(), 7),
_ => panic!("counter missing"),
}
match v.get(&sf) {
Some(FieldValue::OrSet(s)) => assert!(s.contains(b"x")),
_ => panic!("set missing"),
}
match v.get(&rf) {
Some(FieldValue::LwwRegister(r)) => assert_eq!(r.value(), b"hello".to_vec()),
_ => panic!("register missing"),
}
}
#[test]
fn concurrent_remove_loses_to_concurrent_update() {
let a = aid("a");
let b = aid("b");
let f = counter_field("c");
let mut shared = Map::new();
shared.apply(
&a,
&MapOp::Update {
field: f.clone(),
op: NestedOp::Counter(1),
},
);
let mut left = shared.clone();
left.apply(&a, &MapOp::Remove { field: f.clone() });
assert!(!left.contains(&f));
let mut right = shared.clone();
right.apply(
&b,
&MapOp::Update {
field: f.clone(),
op: NestedOp::Counter(2),
},
);
let mut merged = left.clone();
merged.merge(&right);
assert!(
merged.contains(&f),
"add wins on tie: concurrent update beats remove"
);
}
#[test]
fn nested_map_merge_is_recursive() {
let a = aid("a");
let outer_key = map_field("inner");
let inner_counter = counter_field("hits");
let mut left = Map::new();
left.apply(
&a,
&MapOp::Update {
field: outer_key.clone(),
op: NestedOp::Map(Box::new(MapOp::Update {
field: inner_counter.clone(),
op: NestedOp::Counter(3),
})),
},
);
let mut right = Map::new();
right.apply(
&aid("b"),
&MapOp::Update {
field: outer_key.clone(),
op: NestedOp::Map(Box::new(MapOp::Update {
field: inner_counter.clone(),
op: NestedOp::Counter(4),
})),
},
);
left.merge(&right);
match left.get(&outer_key) {
Some(FieldValue::NestedMap(inner)) => match inner.get(&inner_counter) {
Some(FieldValue::Counter(c)) => assert_eq!(c.value(), 7),
_ => panic!("nested counter missing"),
},
_ => panic!("nested map missing"),
}
}
#[test]
fn merge_is_commutative() {
let a = aid("a");
let b = aid("b");
let cf = counter_field("c");
let sf = set_field("s");
let mut x = Map::new();
x.apply(
&a,
&MapOp::Update {
field: cf.clone(),
op: NestedOp::Counter(3),
},
);
x.apply(
&a,
&MapOp::Update {
field: sf.clone(),
op: NestedOp::SetAdd(b"x".to_vec()),
},
);
let mut y = Map::new();
y.apply(
&b,
&MapOp::Update {
field: cf.clone(),
op: NestedOp::Counter(5),
},
);
y.apply(
&b,
&MapOp::Update {
field: sf.clone(),
op: NestedOp::SetAdd(b"y".to_vec()),
},
);
let mut left = x.clone();
left.merge(&y);
let mut right = y.clone();
right.merge(&x);
assert_eq!(left.value(), right.value());
}
#[test]
fn merge_is_idempotent() {
let a = aid("a");
let cf = counter_field("c");
let mut m = Map::new();
m.apply(
&a,
&MapOp::Update {
field: cf,
op: NestedOp::Counter(7),
},
);
let snap = m.clone();
m.merge(&snap);
assert_eq!(m.value(), snap.value());
}
#[test]
fn fields_with_same_name_but_different_types_are_distinct() {
let a = aid("a");
let mut m = Map::new();
let counter_x = counter_field("x");
let flag_x = flag_field("x");
m.apply(
&a,
&MapOp::Update {
field: counter_x.clone(),
op: NestedOp::Counter(1),
},
);
m.apply(
&a,
&MapOp::Update {
field: flag_x.clone(),
op: NestedOp::Flag(true),
},
);
assert_eq!(m.value().len(), 2);
assert!(m.contains(&counter_x));
assert!(m.contains(&flag_x));
}
#[test]
fn field_type_wire_round_trips() {
for ty in [
FieldType::Counter,
FieldType::OrSet,
FieldType::LwwRegister,
FieldType::EwFlag,
FieldType::NestedMap,
] {
assert_eq!(FieldType::from_wire(ty.to_wire()), Some(ty));
}
assert!(FieldType::from_wire(0).is_none());
assert!(FieldType::from_wire(99).is_none());
}
}