use std::borrow::Borrow;
use std::fmt;
use std::hash::BuildHasher;
use std::hash::Hash;
use std::hash::RandomState;
use std::sync::Arc;
use hashbrown::HashTable;
use crate::internal::mutex::Mutex;
use crate::once::OnceCell;
#[cfg(test)]
mod tests;
type Entries<K, V> = HashTable<Arc<Entry<K, V>>>;
struct Entry<K, V> {
hash: u64,
key: K,
cell: OnceCell<V>,
}
enum Lookup<K, V> {
Ready(V),
Pending(Arc<Entry<K, V>>),
}
pub struct OnceMap<K, V, S = RandomState> {
entries: Mutex<Entries<K, V>>,
hasher: S,
}
impl<K, V, S> fmt::Debug for OnceMap<K, V, S> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let (len, pending) = {
let entries = self.entries.lock();
let pending = entries
.iter()
.filter(|entry| !entry.cell.initialized())
.count();
(entries.len(), pending)
};
f.debug_struct("OnceMap")
.field("len", &len)
.field("pending", &pending)
.finish()
}
}
impl<K, V, S> OnceMap<K, V, S>
where
K: Eq + Hash,
S: BuildHasher,
{
fn get_or_insert(&self, key: K) -> Lookup<K, V>
where
V: Clone,
{
let hash = self.hasher.hash_one(&key);
let entry = {
let mut entries = self.entries.lock();
if let Some(entry) = entries
.find(hash, |entry| entry.key.eq(&key))
.map(Arc::clone)
{
entry
} else {
let entry = Arc::new(Entry {
hash,
key,
cell: OnceCell::new(),
});
entries.insert_unique(hash, Arc::clone(&entry), |entry| entry.hash);
entry
}
};
Self::classify(entry)
}
fn classify(entry: Arc<Entry<K, V>>) -> Lookup<K, V>
where
V: Clone,
{
match entry.cell.get().cloned() {
Some(value) => Lookup::Ready(value),
None => Lookup::Pending(entry),
}
}
fn find_entry(
&self,
hash: u64,
matches: impl Fn(&Entry<K, V>) -> bool,
) -> Option<Arc<Entry<K, V>>> {
self.entries
.lock()
.find(hash, |entry| matches(entry))
.cloned()
}
fn get_value<Q>(&self, key: &Q) -> Option<V>
where
K: Borrow<Q>,
Q: Eq + Hash + ?Sized,
V: Clone,
{
let hash = self.hasher.hash_one(key);
let entry = self.find_entry(hash, |entry| entry.key.borrow() == key)?;
entry.cell.get().cloned()
}
fn remove_entry<Q>(&self, key: &Q) -> Option<Arc<Entry<K, V>>>
where
K: Borrow<Q>,
Q: Eq + Hash + ?Sized,
{
let hash = self.hasher.hash_one(key);
let mut entries = self.entries.lock();
let occupied = entries
.find_entry(hash, |entry| entry.key.borrow() == key)
.ok()?;
let (entry, _) = occupied.remove();
drop(entries);
Some(entry)
}
fn cleanup_abandoned_entry(&self, entry: Arc<Entry<K, V>>) {
let removed = {
let mut entries = self.entries.lock();
let Ok(occupied) = entries.find_entry(entry.hash, |stored| Arc::ptr_eq(stored, &entry))
else {
drop(entries);
drop(entry);
return;
};
if Arc::strong_count(&entry) == 2 && !entry.cell.initialized() {
Some(occupied.remove().0)
} else {
drop(entry);
None
}
};
drop(removed);
}
fn insert(&mut self, key: K, value: V) {
let hash = self.hasher.hash_one(&key);
let entry = Arc::new(Entry {
hash,
key,
cell: OnceCell::from_value(value),
});
let mut entries = self.entries.lock();
let replaced = entries
.find_entry(hash, |stored| stored.key.eq(&entry.key))
.ok()
.map(|occupied| occupied.remove().0);
entries.insert_unique(hash, entry, |entry| entry.hash);
drop(entries);
drop(replaced);
}
}
impl<K, V, S> FromIterator<(K, V)> for OnceMap<K, V, S>
where
K: Eq + Hash,
V: Clone,
S: BuildHasher + Default,
{
fn from_iter<T: IntoIterator<Item = (K, V)>>(iter: T) -> Self {
let iter = iter.into_iter();
let mut map = Self {
entries: Mutex::new(HashTable::with_capacity(iter.size_hint().0)),
hasher: S::default(),
};
for (key, value) in iter {
map.insert(key, value);
}
map
}
}
struct ComputeCleanupGuard<'a, K, V, S>
where
K: Eq + Hash,
S: BuildHasher,
{
once_map: &'a OnceMap<K, V, S>,
entry: Option<Arc<Entry<K, V>>>,
}
impl<'a, K, V, S> ComputeCleanupGuard<'a, K, V, S>
where
K: Eq + Hash,
S: BuildHasher,
{
fn new(once_map: &'a OnceMap<K, V, S>, entry: Arc<Entry<K, V>>) -> Self {
Self {
once_map,
entry: Some(entry),
}
}
fn entry(&self) -> &Arc<Entry<K, V>> {
self.entry.as_ref().unwrap()
}
fn dismiss(mut self) {
drop(self.entry.take());
}
}
impl<K, V, S> Drop for ComputeCleanupGuard<'_, K, V, S>
where
K: Eq + Hash,
S: BuildHasher,
{
fn drop(&mut self) {
let Some(entry) = self.entry.take() else {
return;
};
self.once_map.cleanup_abandoned_entry(entry);
}
}
impl<K, V, S> Default for OnceMap<K, V, S>
where
K: Eq + Hash,
V: Clone,
S: BuildHasher + Default,
{
fn default() -> Self {
Self::with_hasher(S::default())
}
}
impl<K, V> OnceMap<K, V, RandomState>
where
K: Eq + Hash,
V: Clone,
{
pub fn new() -> Self {
Self::with_hasher(RandomState::new())
}
}
impl<K, V, S> OnceMap<K, V, S>
where
K: Eq + Hash,
V: Clone,
S: BuildHasher,
{
pub fn with_hasher(hasher: S) -> Self {
Self {
entries: Mutex::new(HashTable::new()),
hasher,
}
}
pub async fn compute<F>(&self, key: K, func: F) -> V
where
F: AsyncFnOnce() -> V,
{
let entry = match self.get_or_insert(key) {
Lookup::Ready(value) => return value,
Lookup::Pending(entry) => entry,
};
let guard = ComputeCleanupGuard::new(self, entry);
let result = guard.entry().cell.get_or_init(func).await.clone();
guard.dismiss();
result
}
pub async fn try_compute<E, F>(&self, key: K, func: F) -> Result<V, E>
where
F: AsyncFnOnce() -> Result<V, E>,
{
let entry = match self.get_or_insert(key) {
Lookup::Ready(value) => return Ok(value),
Lookup::Pending(entry) => entry,
};
let guard = ComputeCleanupGuard::new(self, entry);
let result = guard.entry().cell.get_or_try_init(func).await?.clone();
guard.dismiss();
Ok(result)
}
pub fn get<Q>(&self, key: &Q) -> Option<V>
where
K: Borrow<Q>,
Q: Hash + Eq + ?Sized,
{
self.get_value(key)
}
pub fn discard<Q>(&self, key: &Q)
where
K: Borrow<Q>,
Q: Hash + Eq + ?Sized,
{
drop(self.remove_entry(key));
}
pub fn remove<Q>(&self, key: &Q) -> Option<V>
where
K: Borrow<Q>,
Q: Hash + Eq + ?Sized,
{
let entry = self.remove_entry(key)?;
entry.cell.get().cloned()
}
}