use core::{
hash::Hash,
sync::atomic::{AtomicUsize, Ordering},
};
use dashmap::DashMap;
use std::collections::HashMap;
#[derive(Debug, Clone, PartialEq)]
pub struct Generation(usize);
#[derive(Debug)]
pub(crate) struct Generational<H>(AtomicUsize, H);
impl<H> Generational<H> {
pub fn new(h: H) -> Generational<H> {
Generational(AtomicUsize::new(0), h)
}
pub fn increment_generation(&self) {
self.0.fetch_add(1, Ordering::Release);
}
pub fn get_generation(&self) -> Generation {
Generation(self.0.load(Ordering::Acquire))
}
pub fn get_inner(&self) -> &H {
&self.1
}
}
#[derive(Debug)]
pub struct Registry<K, H>
where
K: Eq + Hash + Clone + 'static,
H: 'static,
{
map: DashMap<K, Generational<H>>,
}
impl<K, H> Registry<K, H>
where
K: Eq + Hash + Clone + 'static,
H: 'static,
{
pub fn new() -> Self {
Self {
map: DashMap::new(),
}
}
pub fn op<I, O, V>(&self, key: K, op: O, init: I) -> V
where
I: FnOnce() -> H,
O: FnOnce(&H) -> V,
{
let valref = self.map.entry(key).or_insert_with(|| {
let value = init();
Generational::new(value)
});
let value = valref.value();
let result = op(value.get_inner());
value.increment_generation();
result
}
pub fn delete(&self, key: &K, generation: Generation) -> bool {
self.map
.remove_if(key, |_, g| g.get_generation() == generation)
.is_some()
}
pub fn get_handles(&self) -> HashMap<K, (Generation, H)>
where
H: Clone,
{
self.collect()
}
pub fn collect<T>(&self) -> T
where
H: Clone,
T: std::iter::FromIterator<(K, (Generation, H))>,
{
self.map_collect(|key, generation, handle| (key.clone(), (generation, handle.clone())))
}
pub fn map_collect<F, R, T>(&self, mut f: F) -> T
where
F: for<'a> FnMut(&'a K, Generation, &'a H) -> R,
T: std::iter::FromIterator<R>,
{
self.map
.iter()
.map(|item| {
let value = item.value();
f(item.key(), value.get_generation(), value.get_inner())
})
.collect()
}
}
impl<K, H> Default for Registry<K, H>
where
K: Eq + Hash + Clone + 'static,
H: 'static,
{
fn default() -> Self {
Registry::new()
}
}
#[cfg(test)]
mod tests {
use super::{Generational, Registry};
use std::sync::{
atomic::{AtomicUsize, Ordering::SeqCst},
Arc,
};
#[test]
fn test_generation() {
let generational = Generational::new(());
let start_gen = generational.get_generation();
let start_gen_extra = generational.get_generation();
assert_eq!(start_gen, start_gen_extra);
generational.increment_generation();
let end_gen = generational.get_generation();
assert_ne!(start_gen, end_gen);
}
#[test]
fn test_registry() {
let registry = Registry::<i32, Arc<AtomicUsize>>::new();
let entries = registry.get_handles();
assert_eq!(entries.len(), 0);
let initial_value = registry.op(
1,
|h| h.fetch_add(1, SeqCst),
|| Arc::new(AtomicUsize::new(42)),
);
assert_eq!(initial_value, 42);
let initial_entries = registry.get_handles();
assert_eq!(initial_entries.len(), 1);
let initial_entry = initial_entries
.into_iter()
.next()
.expect("failed to get first entry");
let (key, (initial_gen, value)) = initial_entry;
assert_eq!(key, 1);
assert_eq!(value.load(SeqCst), 43);
let update_value = registry.op(
1,
|h| h.fetch_add(1, SeqCst),
|| Arc::new(AtomicUsize::new(42)),
);
assert_eq!(update_value, 43);
let updated_entries = registry.get_handles();
assert_eq!(updated_entries.len(), 1);
let updated_entry = updated_entries
.into_iter()
.next()
.expect("failed to get updated entry");
let (key, (updated_gen, value)) = updated_entry;
assert_eq!(key, 1);
assert_eq!(value.load(SeqCst), 44);
assert!(!registry.delete(&key, initial_gen));
let entries = registry.get_handles();
assert_eq!(entries.len(), 1);
assert!(registry.delete(&key, updated_gen));
let entries = registry.get_handles();
assert_eq!(entries.len(), 0);
}
}