use crate::Value;
use rpds::{HashTrieMapSync, RedBlackTreeMapSync};
#[derive(Debug, Clone)]
pub struct PersistentHashSet {
index: HashTrieMapSync<Value, u64>,
order: RedBlackTreeMapSync<u64, Value>,
next_seq: u64,
}
impl PersistentHashSet {
pub fn empty() -> Self {
Self {
index: HashTrieMapSync::new_sync(),
order: RedBlackTreeMapSync::new_sync(),
next_seq: 0,
}
}
pub fn count(&self) -> usize {
self.index.size()
}
pub fn is_empty(&self) -> bool {
self.index.is_empty()
}
pub fn contains(&self, val: &Value) -> bool {
self.index.contains_key(val)
}
pub fn conj(&self, val: Value) -> Self {
if self.index.contains_key(&val) {
return self.clone();
}
let seq = self.next_seq;
Self {
index: self.index.insert(val.clone(), seq),
order: self.order.insert(seq, val),
next_seq: seq + 1,
}
}
pub fn conj_mut(&mut self, val: Value) -> &mut Self {
if !self.index.contains_key(&val) {
let seq = self.next_seq;
self.index.insert_mut(val.clone(), seq);
self.order.insert_mut(seq, val);
self.next_seq = seq + 1;
}
self
}
pub fn disj(&self, val: &Value) -> Self {
match self.index.get(val) {
Some(seq) => Self {
index: self.index.remove(val),
order: self.order.remove(seq),
next_seq: self.next_seq,
},
None => self.clone(),
}
}
pub fn iter(&self) -> impl Iterator<Item = &Value> {
self.order.iter().map(|(_, v)| v)
}
}
impl FromIterator<Value> for PersistentHashSet {
fn from_iter<I: IntoIterator<Item = Value>>(iter: I) -> Self {
let mut s = PersistentHashSet::empty();
for v in iter {
s.conj_mut(v);
}
s
}
}
impl PartialEq for PersistentHashSet {
fn eq(&self, other: &Self) -> bool {
if self.count() != other.count() {
return false;
}
self.iter().all(|k| other.contains(k))
}
}
impl cljrs_gc::Trace for PersistentHashSet {
fn trace(&self, visitor: &mut cljrs_gc::MarkVisitor) {
for (_, v) in self.order.iter() {
v.trace(visitor);
}
}
fn gc_size_extra(&self) -> usize {
let n = self.index.size();
n * (88 + std::mem::size_of::<Value>())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Value;
fn int(n: i64) -> Value {
Value::Long(n)
}
#[test]
fn test_basic() {
let s = PersistentHashSet::empty();
let s = s.conj(int(1)).conj(int(2)).conj(int(3));
assert_eq!(s.count(), 3);
assert!(s.contains(&int(1)));
assert!(s.contains(&int(2)));
assert!(!s.contains(&int(99)));
}
#[test]
fn test_idempotent_conj() {
let s = PersistentHashSet::empty().conj(int(1)).conj(int(1));
assert_eq!(s.count(), 1);
}
#[test]
fn test_disj() {
let s = PersistentHashSet::empty().conj(int(1)).conj(int(2));
let s2 = s.disj(&int(1));
assert!(!s2.contains(&int(1)));
assert!(s2.contains(&int(2)));
assert_eq!(s2.count(), 1);
}
#[test]
fn test_equality_order_independent() {
let a = PersistentHashSet::from_iter([int(1), int(2), int(3)]);
let b = PersistentHashSet::from_iter([int(3), int(1), int(2)]);
assert_eq!(a, b);
}
#[test]
fn test_iteration_order_matches_insertion_order() {
let s = PersistentHashSet::empty()
.conj(int(10))
.conj(int(3))
.conj(int(7))
.conj(int(1));
let iterated: Vec<Value> = s.iter().cloned().collect();
assert_eq!(iterated, vec![int(10), int(3), int(7), int(1)]);
}
}