use crate::containers::HashSet;
use crate::containers::StorageKey;
use crate::pumpkin_assert_moderate;
use crate::pumpkin_assert_simple;
#[derive(Debug, Clone)]
pub struct SparseSet<T> {
size: usize,
domain: Vec<T>,
indices: Vec<usize>,
mapping: fn(&T) -> i32,
index_offset: i32,
}
impl<T: StorageKey> SparseSet<T> {
pub fn new(input: Vec<T>) -> Self {
Self::new_with_mapping(input, |element: &T| element.index() as i32)
}
}
impl<T> SparseSet<T> {
pub fn new_with_mapping(input: Vec<T>, mapping: fn(&T) -> i32) -> Self {
let input_len = input.len();
let mut min_index = 0;
let mut max_index = 0;
let mut used_indices = HashSet::new();
for element in input.iter() {
let index = (mapping)(element);
let not_previously_inserted = used_indices.insert(index);
pumpkin_assert_simple!(
not_previously_inserted,
"Two elements in the provided `input` map to the same index."
);
min_index = min_index.min(index);
max_index = max_index.max(index);
}
pumpkin_assert_simple!(min_index <= max_index);
let mut indices =
std::iter::repeat_n(usize::MAX, (max_index.abs() + min_index.abs()) as usize + 1)
.collect::<Vec<_>>();
for (i, element) in input.iter().enumerate().collect::<Vec<_>>() {
indices[((mapping)(element) - min_index) as usize] = i;
}
SparseSet {
size: input_len,
domain: input,
indices,
mapping,
index_offset: -min_index,
}
}
fn get_mapping(&self, element: &T) -> usize {
let output_index = (self.mapping)(element) + self.index_offset;
output_index.try_into().unwrap()
}
pub fn set_to_empty(&mut self) {
self.indices = vec![usize::MAX; self.indices.len()];
self.domain.clear();
self.size = 0;
}
pub fn restore_temporarily_removed(&mut self) {
self.size = self.domain.len();
}
pub fn is_empty(&self) -> bool {
self.size == 0
}
pub fn len(&self) -> usize {
self.size
}
pub fn get(&self, index: usize) -> &T {
pumpkin_assert_simple!(index < self.size);
&self.domain[index]
}
fn swap(&mut self, i: usize, j: usize) {
self.domain.swap(i, j);
let index_i = self.get_mapping(&self.domain[i]);
self.indices[index_i] = i;
let index_j = self.get_mapping(&self.domain[j]);
self.indices[index_j] = j;
}
pub fn remove(&mut self, to_remove: &T) {
if self.indices[self.get_mapping(to_remove)] < self.size {
self.size -= 1;
if self.size > 0 {
self.swap(self.indices[self.get_mapping(to_remove)], self.size);
}
self.swap(
self.indices[self.get_mapping(to_remove)],
self.domain.len() - 1,
);
let element = self.domain.pop().expect("Has to have something to pop.");
pumpkin_assert_moderate!((self.mapping)(&element) == (self.mapping)(to_remove));
let to_remove_index = self.get_mapping(to_remove);
self.indices[to_remove_index] = usize::MAX;
} else if self.indices[self.get_mapping(to_remove)] < self.domain.len() {
self.swap(
self.indices[self.get_mapping(to_remove)],
self.domain.len() - 1,
);
let element = self.domain.pop().expect("Has to have something to pop.");
pumpkin_assert_moderate!((self.mapping)(&element) == (self.mapping)(to_remove));
let to_remove_index = self.get_mapping(to_remove);
self.indices[to_remove_index] = usize::MAX;
}
}
pub fn remove_temporarily(&mut self, to_remove: &T) {
if self.indices[self.get_mapping(to_remove)] < self.size {
self.size -= 1;
self.swap(self.indices[self.get_mapping(to_remove)], self.size);
}
}
pub fn contains(&self, element: &T) -> bool {
self.get_mapping(element) < self.indices.len()
&& self.indices[self.get_mapping(element)] < self.size
}
pub fn accommodate(&mut self, element: &T) {
let index = self.get_mapping(element);
if self.indices.len() <= index {
self.indices.resize(index + 1, usize::MAX);
}
}
pub fn insert(&mut self, element: T) {
if !self.contains(&element) {
self.accommodate(&element);
let mut index = self.indices[self.get_mapping(&element)];
if index >= self.domain.len() {
index = self.domain.len();
let element_index = self.get_mapping(&element);
self.indices[element_index] = index;
self.domain.push(element);
}
self.swap(self.size, index);
self.size += 1;
}
}
pub fn iter(&self) -> impl Iterator<Item = &T> {
self.domain[..self.size].iter()
}
pub fn out_of_domain(&self) -> impl Iterator<Item = &T> {
self.domain[self.size..].iter()
}
}
impl<T> IntoIterator for SparseSet<T> {
type Item = T;
type IntoIter = std::iter::Take<std::vec::IntoIter<T>>;
fn into_iter(self) -> Self::IntoIter {
self.domain.into_iter().take(self.size)
}
}
#[cfg(test)]
mod tests {
use super::SparseSet;
fn mapping_function(input: &i32) -> i32 {
*input
}
#[test]
fn test_len() {
let sparse_set = SparseSet::new_with_mapping(vec![0, 1, 2], mapping_function);
assert_eq!(sparse_set.len(), 3);
}
#[test]
fn removal() {
let mut sparse_set = SparseSet::new_with_mapping(vec![0, 1, 2], mapping_function);
sparse_set.remove(&1);
assert_eq!(sparse_set.domain, vec![0, 2]);
assert_eq!(sparse_set.size, 2);
assert_eq!(sparse_set.indices, vec![0, usize::MAX, 1]);
}
#[test]
fn removal_adjusts_size() {
let mut sparse_set = SparseSet::new_with_mapping(vec![0, 1, 2], mapping_function);
assert_eq!(sparse_set.size, 3);
sparse_set.remove(&0);
assert_eq!(sparse_set.size, 2);
}
#[test]
fn remove_all_elements_leads_to_empty_set() {
let mut sparse_set = SparseSet::new_with_mapping(vec![0, 1, 2], mapping_function);
sparse_set.remove(&0);
sparse_set.remove(&1);
sparse_set.remove(&2);
assert!(sparse_set.is_empty());
}
#[test]
fn iter1() {
let sparse_set = SparseSet::new_with_mapping(vec![5, 10, 2], mapping_function);
let v: Vec<i32> = sparse_set.iter().copied().collect();
assert_eq!(v.len(), 3);
assert!(v.contains(&10));
assert!(v.contains(&5));
assert!(v.contains(&2));
}
#[test]
fn iter2() {
let mut sparse_set = SparseSet::new_with_mapping(vec![5, 10, 2], mapping_function); sparse_set.insert(100); sparse_set.insert(2); sparse_set.insert(20); sparse_set.remove(&10); sparse_set.insert(10); sparse_set.remove(&10);
let v: Vec<i32> = sparse_set.iter().copied().collect();
assert_eq!(v.len(), 4);
assert!(v.contains(&5));
assert!(v.contains(&2));
assert!(v.contains(&100));
assert!(v.contains(&20));
assert!(!v.contains(&10));
}
#[test]
fn remove_temporarily_simple() {
let mut sparse_set = SparseSet::new_with_mapping(vec![0], mapping_function);
sparse_set.remove_temporarily(&0);
sparse_set.insert(0);
sparse_set.remove_temporarily(&0);
assert!(sparse_set.is_empty())
}
#[test]
fn remove_temporarily() {
let mut sparse_set = SparseSet::new_with_mapping(vec![2, 0, 1], mapping_function);
assert!(!sparse_set.is_empty());
sparse_set.remove_temporarily(&0);
sparse_set.insert(0);
sparse_set.remove_temporarily(&0);
assert!(!sparse_set.contains(&0));
assert!(!sparse_set.is_empty());
sparse_set.remove_temporarily(&0);
assert!(!sparse_set.contains(&0));
sparse_set.remove_temporarily(&0);
sparse_set.remove_temporarily(&2);
assert!(!sparse_set.contains(&2));
sparse_set.remove_temporarily(&1);
assert!(!sparse_set.contains(&1));
assert!(sparse_set.is_empty());
sparse_set.insert(1);
assert!(sparse_set.contains(&1));
assert!(!sparse_set.contains(&0));
assert!(sparse_set.contains(&1));
assert!(!sparse_set.contains(&2));
sparse_set.restore_temporarily_removed();
assert!(sparse_set.contains(&0));
assert!(sparse_set.contains(&1));
assert!(sparse_set.contains(&2));
}
#[test]
fn remove_temporarily_non_continuous() {
let mut sparse_set = SparseSet::new_with_mapping(vec![5, 10, 2], mapping_function);
sparse_set.remove_temporarily(&10);
assert!(!sparse_set.contains(&10));
sparse_set.remove_temporarily(&5);
sparse_set.remove_temporarily(&2);
assert!(sparse_set.is_empty());
}
#[test]
fn remove_temporarily_non_continuous_spanning() {
let mut sparse_set = SparseSet::new_with_mapping(vec![5, 10, -2], mapping_function);
sparse_set.remove_temporarily(&10);
assert!(!sparse_set.contains(&10));
sparse_set.remove_temporarily(&-2);
assert!(!sparse_set.contains(&-2));
sparse_set.remove_temporarily(&5);
assert!(sparse_set.is_empty());
}
#[test]
fn remove_temporarily_non_continuous_negative() {
let mut sparse_set = SparseSet::new_with_mapping(vec![-5, -10, -2], mapping_function);
sparse_set.remove_temporarily(&-10);
assert!(!sparse_set.contains(&-10));
sparse_set.remove_temporarily(&-2);
assert!(!sparse_set.contains(&-2));
sparse_set.remove_temporarily(&-5);
assert!(sparse_set.is_empty());
}
}