cljrs_value/collections/
hash_set.rs1use crate::Value;
2use rpds::{HashTrieMapSync, RedBlackTreeMapSync};
3
4#[derive(Debug, Clone)]
14pub struct PersistentHashSet {
15 index: HashTrieMapSync<Value, u64>,
16 order: RedBlackTreeMapSync<u64, Value>,
17 next_seq: u64,
18}
19
20impl PersistentHashSet {
21 pub fn empty() -> Self {
22 Self {
23 index: HashTrieMapSync::new_sync(),
24 order: RedBlackTreeMapSync::new_sync(),
25 next_seq: 0,
26 }
27 }
28
29 pub fn count(&self) -> usize {
30 self.index.size()
31 }
32
33 pub fn is_empty(&self) -> bool {
34 self.index.is_empty()
35 }
36
37 pub fn contains(&self, val: &Value) -> bool {
38 self.index.contains_key(val)
39 }
40
41 pub fn conj(&self, val: Value) -> Self {
45 if self.index.contains_key(&val) {
46 return self.clone();
47 }
48 let seq = self.next_seq;
49 Self {
50 index: self.index.insert(val.clone(), seq),
51 order: self.order.insert(seq, val),
52 next_seq: seq + 1,
53 }
54 }
55
56 pub fn conj_mut(&mut self, val: Value) -> &mut Self {
57 if !self.index.contains_key(&val) {
58 let seq = self.next_seq;
59 self.index.insert_mut(val.clone(), seq);
60 self.order.insert_mut(seq, val);
61 self.next_seq = seq + 1;
62 }
63 self
64 }
65
66 pub fn disj(&self, val: &Value) -> Self {
68 match self.index.get(val) {
69 Some(seq) => Self {
70 index: self.index.remove(val),
71 order: self.order.remove(seq),
72 next_seq: self.next_seq,
73 },
74 None => self.clone(),
75 }
76 }
77
78 pub fn iter(&self) -> impl Iterator<Item = &Value> {
80 self.order.iter().map(|(_, v)| v)
81 }
82}
83
84impl FromIterator<Value> for PersistentHashSet {
85 fn from_iter<I: IntoIterator<Item = Value>>(iter: I) -> Self {
86 let mut s = PersistentHashSet::empty();
87 for v in iter {
88 s.conj_mut(v);
89 }
90 s
91 }
92}
93
94impl PartialEq for PersistentHashSet {
95 fn eq(&self, other: &Self) -> bool {
96 if self.count() != other.count() {
97 return false;
98 }
99 self.iter().all(|k| other.contains(k))
100 }
101}
102
103impl cljrs_gc::Trace for PersistentHashSet {
104 fn trace(&self, visitor: &mut cljrs_gc::MarkVisitor) {
105 for (_, v) in self.order.iter() {
106 v.trace(visitor);
107 }
108 }
109
110 fn gc_size_extra(&self) -> usize {
111 let n = self.index.size();
114 n * (88 + std::mem::size_of::<Value>())
115 }
116}
117
118#[cfg(test)]
119mod tests {
120 use super::*;
121 use crate::Value;
122
123 fn int(n: i64) -> Value {
124 Value::Long(n)
125 }
126
127 #[test]
128 fn test_basic() {
129 let s = PersistentHashSet::empty();
130 let s = s.conj(int(1)).conj(int(2)).conj(int(3));
131 assert_eq!(s.count(), 3);
132 assert!(s.contains(&int(1)));
133 assert!(s.contains(&int(2)));
134 assert!(!s.contains(&int(99)));
135 }
136
137 #[test]
138 fn test_idempotent_conj() {
139 let s = PersistentHashSet::empty().conj(int(1)).conj(int(1));
140 assert_eq!(s.count(), 1);
141 }
142
143 #[test]
144 fn test_disj() {
145 let s = PersistentHashSet::empty().conj(int(1)).conj(int(2));
146 let s2 = s.disj(&int(1));
147 assert!(!s2.contains(&int(1)));
148 assert!(s2.contains(&int(2)));
149 assert_eq!(s2.count(), 1);
150 }
151
152 #[test]
153 fn test_equality_order_independent() {
154 let a = PersistentHashSet::from_iter([int(1), int(2), int(3)]);
155 let b = PersistentHashSet::from_iter([int(3), int(1), int(2)]);
156 assert_eq!(a, b);
157 }
158
159 #[test]
160 fn test_iteration_order_matches_insertion_order() {
161 let s = PersistentHashSet::empty()
162 .conj(int(10))
163 .conj(int(3))
164 .conj(int(7))
165 .conj(int(1));
166 let iterated: Vec<Value> = s.iter().cloned().collect();
167 assert_eq!(iterated, vec![int(10), int(3), int(7), int(1)]);
168 }
169}