Skip to main content

causal_length/
set.rs

1use super::*;
2use crate::register::Register;
3use std::borrow::Borrow;
4use std::cmp::max;
5use std::collections::HashMap;
6use std::hash::Hash;
7
8#[derive(Clone, Debug, Eq, PartialEq)]
9struct SubRegister<Tag, CL>
10where
11    Tag: TagT,
12    CL: CausalLength,
13{
14    tag: Tag,
15    length: CL,
16}
17
18/// Causal Length Set
19///
20/// Set implements the set described in the paper, with the addition of a tag. Set only uses the
21/// tag for garbage collection of old removed members.
22#[derive(Clone, Debug, Default, PartialEq, Eq)]
23pub struct Set<T, Tag, CL>
24where
25    T: Key,
26    Tag: TagT,
27    CL: CausalLength,
28{
29    // HashMap, because the "set" needs to allow mutating the tag and causal length.
30    map: HashMap<T, SubRegister<Tag, CL>>,
31}
32
33impl<T, Tag, CL> Set<T, Tag, CL>
34where
35    T: Key,
36    Tag: TagT,
37    CL: CausalLength,
38{
39    /// Create a new empty `Set`
40    pub fn new() -> Set<T, Tag, CL> {
41        Set {
42            map: HashMap::new(),
43        }
44    }
45
46    /// Returns `None` if `member` is not present in the set. If present returns `Some(Tag)`
47    pub fn get<Q>(&self, member: Q) -> Option<Tag>
48    where
49        Q: Borrow<T>,
50    {
51        if let Some(e) = self.map.get(member.borrow()).to_owned() {
52            if e.length.is_odd() {
53                return Some(e.tag);
54            }
55        }
56        None
57    }
58
59    /// Returns true if the set contains a value.
60    pub fn contains<Q>(&self, member: Q) -> bool
61    where
62        Q: Borrow<T>,
63    {
64        self.get(member).is_some()
65    }
66
67    /// Add a value to a set.
68    pub fn add(&mut self, member: T, tag: Tag) {
69        let one: CL = CL::one();
70        let mut e = self
71            .map
72            .entry(member)
73            .or_insert(SubRegister { tag, length: one });
74        // s{e |-> s(e)+1} if even
75        //s if odd s(e)
76        if e.length.is_even() {
77            e.length = e.length + one;
78        }
79        // always use the max value of tag
80        e.tag = max(e.tag, tag);
81    }
82
83    /// Removes a value from the set.
84    pub fn remove(&mut self, member: T, tag: Tag) {
85        self.map.entry(member).and_modify(|e| {
86            // {} if even(s(e))
87            // { e |-> s(e) + 1 } if odd(s(e))
88            if e.length.is_odd() {
89                e.length = e.length + CL::one()
90            }
91            e.tag = max(e.tag, tag);
92        });
93        // ignore attempts to remove items that aren't present...
94    }
95
96    /// An iterator visiting all elements and tags in arbitrary order.
97    pub fn iter(&self) -> impl Iterator<Item = (&T, Tag)> + '_ {
98        self.map
99            .iter()
100            .filter(|(_k, v)| v.length.is_odd())
101            .map(|(k, v)| (k, v.tag))
102    }
103
104    /// An iterator visiting all registers in arbitrary order.
105    pub fn register_iter(&self) -> impl Iterator<Item = Register<T, Tag, CL>> + '_ {
106        self.map.iter().map(|(k, v)| Register {
107            item: k.clone(),
108            tag: v.tag,
109            length: v.length,
110        })
111    }
112
113    /// Merge a delta [Register] into a set.
114    ///
115    /// Remove registers with a tag value less than `min_tag` will be ignored.
116    pub fn merge_register(&mut self, delta: Register<T, Tag, CL>, min_tag: Tag) {
117        if delta.length.is_even() && delta.tag < min_tag {
118            // ignore excessively old remove records
119            return;
120        }
121        let Register { item, tag, length } = delta;
122        match self.map.entry(item) {
123            Entry::Occupied(mut e) => {
124                let e = e.get_mut();
125                // (s⊔s′)(e) = max(s(e),s′(e))
126                e.tag = max(e.tag, tag);
127                e.length = max(e.length, length);
128            }
129            Entry::Vacant(e) => {
130                e.insert(SubRegister { tag, length });
131            }
132        }
133    }
134
135    /// Merge two sets.
136    ///
137    /// Remove deltas with a tag value less than `min_tag` will be ignored.
138    pub fn merge(&mut self, other: &Self, min_tag: Tag) {
139        for delta in other.register_iter() {
140            self.merge_register(delta, min_tag);
141        }
142    }
143
144    /// Filter out old remove tombstone deltas from the set.
145    ///
146    /// Remove deltas with a tag value less than `min_tag` will be removed.
147    pub fn retain(&mut self, min_tag: Tag) {
148        self.map
149            .retain(|_k, SubRegister { tag, length }| length.is_odd() || min_tag < *tag);
150    }
151}
152
153#[cfg(feature = "serialization")]
154mod serialization {
155    use super::*;
156    use serde::de::{SeqAccess, Visitor};
157    use serde::ser::SerializeSeq;
158    use serde::{Deserialize, Deserializer, Serialize, Serializer};
159    use std::fmt::Formatter;
160    use std::marker::PhantomData;
161
162    impl<T, Tag, CL> Serialize for Set<T, Tag, CL>
163    where
164        T: Key + Serialize,
165        Tag: TagT + Serialize,
166        CL: CausalLength + Serialize,
167    {
168        fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
169        where
170            S: Serializer,
171        {
172            let mut seq = serializer.serialize_seq(Some(self.map.len()))?;
173            for member in self.register_iter() {
174                seq.serialize_element(&(member.item, member.tag, member.length))?;
175            }
176            seq.end()
177        }
178    }
179
180    struct DeltaVisitor<T, Tag, CL>(PhantomData<T>, PhantomData<Tag>, PhantomData<CL>);
181
182    impl<'de, T, Tag, CL> Visitor<'de> for DeltaVisitor<T, Tag, CL>
183    where
184        T: Key + Deserialize<'de>,
185        Tag: TagT + Deserialize<'de>,
186        CL: CausalLength + Deserialize<'de>,
187    {
188        type Value = HashMap<T, SubRegister<Tag, CL>>;
189
190        fn expecting(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
191            formatter.write_str("a tuple of key, value, tag, and causal length")
192        }
193
194        fn visit_seq<A>(self, mut seq: A) -> std::result::Result<Self::Value, A::Error>
195        where
196            A: SeqAccess<'de>,
197        {
198            let mut map: HashMap<T, SubRegister<Tag, CL>> =
199                HashMap::with_capacity(seq.size_hint().unwrap_or(0));
200            while let Some(d) = seq.next_element::<(T, Tag, CL)>()? {
201                map.insert(
202                    d.0,
203                    SubRegister {
204                        tag: d.1,
205                        length: d.2,
206                    },
207                );
208            }
209            Ok(map)
210        }
211    }
212
213    impl<'de, T, Tag, CL> Deserialize<'de> for Set<T, Tag, CL>
214    where
215        T: Eq + Hash + Clone + Deserialize<'de>,
216        Tag: TagT + Deserialize<'de>,
217        CL: CausalLength + Deserialize<'de>,
218    {
219        fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
220        where
221            D: Deserializer<'de>,
222        {
223            let visitor = DeltaVisitor::<T, Tag, CL>(PhantomData, PhantomData, PhantomData);
224            let map = deserializer.deserialize_seq(visitor)?;
225
226            Ok(Set { map })
227        }
228    }
229}
230
231#[cfg(feature = "serialization")]
232pub use serialization::*;
233use std::collections::hash_map::Entry;
234
235#[cfg(test)]
236mod tests {
237    use super::*;
238    use quickcheck_macros::quickcheck;
239    use rand::seq::SliceRandom;
240
241    #[test]
242    fn test_add() {
243        let later_time = 1;
244        let mut cls: Set<&str, u32, u16> = Set::new();
245
246        cls.add("foo", later_time);
247        cls.add("foo", later_time);
248        cls.add("foo", later_time);
249        assert_eq!(cls.map.len(), 1);
250        assert_eq!(
251            cls.map.get("foo"),
252            Some(&SubRegister {
253                tag: later_time,
254                length: 1
255            })
256        );
257        assert_eq!(cls.contains("foo"), true);
258        assert_eq!(cls.get("bar"), None);
259    }
260
261    #[test]
262    fn test_remove() {
263        let time_1 = 1;
264        let time_2 = 2;
265        let time_3 = 3;
266        let mut cls: Set<&str, u32, u16> = Set::new();
267
268        cls.add("foo", time_1);
269        cls.add("bar", time_1);
270        cls.remove("foo", time_2);
271        cls.remove("bar", time_2);
272        cls.add("bar", time_3);
273        // check map
274        assert_eq!(cls.map.len(), 2);
275        assert_eq!(
276            cls.map.get(&"bar"),
277            Some(&SubRegister {
278                tag: time_3,
279                length: 3
280            })
281        );
282        assert_eq!(
283            cls.map.get(&"foo"),
284            Some(&SubRegister {
285                tag: time_2,
286                length: 2
287            })
288        );
289        // check edges
290        let values: Vec<(&&str, u32)> = cls.iter().collect();
291        assert_eq!(values.len(), 1);
292        assert_eq!(values[0], (&"bar", time_3));
293    }
294
295    #[test]
296    fn test_merge() {
297        let time_0 = 0;
298        let time_1 = 1;
299        let time_2 = 2;
300        let time_3 = 3;
301        let mut cls1: Set<&str, u32, u16> = Set::new();
302        let mut cls2: Set<&str, u32, u16> = Set::new();
303
304        cls1.add("foo", time_1);
305        cls1.add("bar", time_1);
306        cls2.merge(&cls1, time_0);
307        cls2.remove("foo", time_2);
308        cls1.remove("bar", time_2);
309        cls1.remove("bar", time_2);
310        cls2.merge(&cls1, time_0);
311        cls2.add("bar", time_3);
312        // check map
313        assert_eq!(cls2.map.len(), 2);
314        assert_eq!(
315            cls2.map.get(&"bar"),
316            Some(&SubRegister {
317                tag: time_3,
318                length: 3
319            })
320        );
321        assert_eq!(
322            cls2.map.get(&"foo"),
323            Some(&SubRegister {
324                tag: time_2,
325                length: 2
326            })
327        );
328        // check edges
329        let values: Vec<(&&str, u32)> = cls2.iter().collect();
330        assert_eq!(values.len(), 1);
331        assert_eq!(values[0], (&"bar", time_3));
332    }
333
334    #[test]
335    fn test_retain() {
336        let time_0 = 0;
337        let time_1 = 1;
338        let time_2 = 2;
339        let time_3 = 3;
340        let mut cls: Set<&str, u32, u16> = Set::new();
341
342        cls.add("foo", time_0);
343        cls.add("bar", time_0);
344        cls.remove("foo", time_1);
345        cls.remove("bar", time_1);
346        cls.add("bar", time_2);
347        // check map
348        assert_eq!(cls.map.len(), 2);
349        assert_eq!(
350            cls.map.get(&"bar"),
351            Some(&SubRegister {
352                tag: time_2,
353                length: 3
354            })
355        );
356        assert_eq!(
357            cls.map.get(&"foo"),
358            Some(&SubRegister {
359                tag: time_1,
360                length: 2
361            })
362        );
363        // check edges
364        let values: Vec<(&&str, u32)> = cls.iter().collect();
365        assert_eq!(values.len(), 1);
366        assert_eq!(values[0], (&"bar", time_2));
367        // now clear old removes
368        cls.retain(time_3);
369        assert_eq!(cls.map.len(), 1);
370        assert_eq!(
371            cls.map.get(&"bar"),
372            Some(&SubRegister {
373                tag: time_2,
374                length: 3
375            })
376        );
377        // attempt to merge an out of date remove
378        cls.merge_register(
379            Register {
380                item: &"bar",
381                tag: time_2,
382                length: 2,
383            },
384            time_0,
385        );
386        assert_eq!(cls.map.len(), 1);
387        assert_eq!(
388            cls.map.get(&"bar"),
389            Some(&SubRegister {
390                tag: time_2,
391                length: 3
392            })
393        );
394    }
395
396    #[cfg(feature = "serialization")]
397    #[test]
398    fn test_serialization() {
399        let time_1 = 1;
400        let time_2 = 2;
401        let time_3 = 3;
402        let mut cls: Set<&str, u32, u16> = Set::new();
403
404        cls.add("foo", time_1);
405        cls.add("bar", time_1);
406        cls.remove("foo", time_2);
407        cls.remove("bar", time_2);
408        cls.add("bar", time_3);
409
410        let data = serde_json::to_vec(&cls).unwrap();
411        let cls2: Set<&str, u32, u16> = serde_json::from_slice(&data).unwrap();
412        assert_eq!(cls.map, cls2.map);
413    }
414
415    #[test]
416    fn test_order_independence() {
417        let mut m: Set<&str, u32, u16> = Set::new();
418        let mut v: Vec<Register<&str, u32, u16>> = vec![];
419
420        for i in 0..1000 {
421            v.push(Register {
422                item: "foo",
423                tag: i as u32,
424                length: i as u16,
425            });
426        }
427
428        // now randomize the updates
429        v.shuffle(&mut rand::thread_rng());
430
431        for r in v {
432            m.merge_register(r, 0);
433        }
434        assert_eq!(
435            m.map.get("foo"),
436            Some(&SubRegister {
437                tag: 999,
438                length: 999
439            })
440        );
441    }
442
443    fn merge(mut acc: Set<u8, u8, u8>, el: &Register<u8, u8, u8>) -> Set<u8, u8, u8> {
444        acc.merge_register(el.clone(), 0);
445        acc
446    }
447
448    #[quickcheck]
449    fn is_merge_commutative(xs: Vec<Register<u8, u8, u8>>) -> bool {
450        let left = xs.iter().fold(Set::default(), merge);
451        let right = xs.iter().rfold(Set::default(), merge);
452        left == right
453    }
454
455    #[quickcheck]
456    fn is_merge_order_independent(xs: Vec<Register<u8, u8, u8>>) -> bool {
457        let mut copy = xs.clone();
458        copy.shuffle(&mut rand::thread_rng());
459        let left = xs.iter().fold(Set::default(), merge);
460        let right = copy.iter().rfold(Set::default(), merge);
461        left == right
462    }
463
464    use quickcheck::{Arbitrary, Gen};
465    #[derive(Clone, Debug)]
466    enum Op {
467        Insert(u8),
468        Get(u8),
469        Delete(u8),
470    }
471
472    const KEY_SPACE: u8 = 20;
473
474    impl Arbitrary for Op {
475        fn arbitrary(g: &mut Gen) -> Op {
476            let k: u8 = u8::arbitrary(g) % KEY_SPACE;
477            let n: u8 = u8::arbitrary(g) % 4;
478
479            match n {
480                0 => Op::Insert(k),
481                1 => Op::Delete(k),
482                2 | 3 => Op::Get(k),
483                _ => Op::Get(k),
484            }
485        }
486    }
487
488    #[quickcheck]
489    fn implementation_matches_model(ops: Vec<Op>) -> bool {
490        let mut implementation: Set<u8, u8, u8> = Set::new();
491        let mut model = std::collections::HashSet::new();
492
493        for op in ops {
494            match op {
495                Op::Insert(k) => {
496                    implementation.add(k, 0);
497                    model.insert(k);
498                }
499                Op::Get(k) => {
500                    if implementation.get(&k).is_some() != model.get(&k).is_some() {
501                        return false;
502                    }
503                }
504                Op::Delete(k) => {
505                    implementation.remove(k, 0);
506                    model.remove(&k);
507                }
508            }
509        }
510
511        true
512    }
513}