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#[derive(Clone, Debug, Default, PartialEq, Eq)]
23pub struct Set<T, Tag, CL>
24where
25 T: Key,
26 Tag: TagT,
27 CL: CausalLength,
28{
29 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 pub fn new() -> Set<T, Tag, CL> {
41 Set {
42 map: HashMap::new(),
43 }
44 }
45
46 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 pub fn contains<Q>(&self, member: Q) -> bool
61 where
62 Q: Borrow<T>,
63 {
64 self.get(member).is_some()
65 }
66
67 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 if e.length.is_even() {
77 e.length = e.length + one;
78 }
79 e.tag = max(e.tag, tag);
81 }
82
83 pub fn remove(&mut self, member: T, tag: Tag) {
85 self.map.entry(member).and_modify(|e| {
86 if e.length.is_odd() {
89 e.length = e.length + CL::one()
90 }
91 e.tag = max(e.tag, tag);
92 });
93 }
95
96 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 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 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 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 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 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 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 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 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 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 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 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 let values: Vec<(&&str, u32)> = cls.iter().collect();
365 assert_eq!(values.len(), 1);
366 assert_eq!(values[0], (&"bar", time_2));
367 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 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 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}