use std::borrow::Borrow;
use std::cell::UnsafeCell;
use std::collections::hash_map::{self, Entry, RandomState};
use std::collections::HashMap;
use std::convert::Infallible;
use std::hash::{BuildHasher, Hash};
use std::marker::PhantomPinned;
use std::mem::{transmute, ManuallyDrop};
use std::pin::Pin;
use std::sync::{Mutex, MutexGuard};
macro_rules! lock {
($mutex:expr) => {
match $mutex.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
}
};
}
macro_rules! get_mut {
(let $target:ident, $mutex:expr) => {
let mut $target = $mutex.get_mut();
let $target = match $target {
Ok(guard) => guard,
Err(ref mut poisoned) => poisoned.get_mut(),
};
};
}
struct StableEntryInner<K, V> {
key: K,
value: UnsafeCell<V>,
_pin: PhantomPinned,
}
struct StableEntry<K, V>(Pin<Box<StableEntryInner<K, V>>>);
impl<K, V> StableEntry<K, V> {
fn new(key: K, value: V) -> Self {
StableEntry(Box::pin(StableEntryInner {
key,
value: UnsafeCell::new(value),
_pin: PhantomPinned,
}))
}
fn key(&self) -> &K {
&self.0.key
}
fn value(&self) -> &V {
unsafe { &*self.0.value.get() }
}
fn value_ptr(&self) -> *mut V {
self.0.value.get()
}
fn into_value(self) -> V {
unsafe { Pin::into_inner_unchecked(self.0) }
.value
.into_inner()
}
}
impl<K: Clone, V: Clone> Clone for StableEntry<K, V> {
fn clone(&self) -> Self {
StableEntry::new(self.key().clone(), self.value().clone())
}
}
impl<K: Hash, V> Hash for StableEntry<K, V> {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.key().hash(state);
}
}
impl<K: PartialEq, V> PartialEq for StableEntry<K, V> {
fn eq(&self, other: &Self) -> bool {
self.key().eq(other.key())
}
}
impl<K: Eq, V> Eq for StableEntry<K, V> {}
impl<K, V, Q: ?Sized> Borrow<BorrowedKey<Q>> for StableEntry<K, V>
where
K: Borrow<Q>,
{
fn borrow(&self) -> &BorrowedKey<Q> {
BorrowedKey::from_ref(self.key().borrow())
}
}
#[repr(transparent)]
struct BorrowedKey<Q: ?Sized>(Q);
impl<Q: ?Sized> BorrowedKey<Q> {
fn from_ref(key: &Q) -> &Self {
unsafe { &*(key as *const Q as *const BorrowedKey<Q>) }
}
}
impl<Q: Hash + ?Sized> Hash for BorrowedKey<Q> {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.0.hash(state);
}
}
impl<Q: PartialEq + ?Sized> PartialEq for BorrowedKey<Q> {
fn eq(&self, other: &Self) -> bool {
self.0.eq(&other.0)
}
}
impl<Q: Eq + ?Sized> Eq for BorrowedKey<Q> {}
type InnerMap<K, V, S> = HashMap<StableEntry<K, V>, (), S>;
struct DebugMap<'a, K, V, S>(&'a InnerMap<K, V, S>);
impl<K: std::fmt::Debug, V: std::fmt::Debug, S> std::fmt::Debug for DebugMap<'_, K, V, S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut map = f.debug_map();
for entry in self.0.keys() {
map.entry(entry.key(), entry.value());
}
map.finish()
}
}
pub struct MemoMap<K, V, S = RandomState> {
inner: Mutex<InnerMap<K, V, S>>,
}
impl<K: std::fmt::Debug, V: std::fmt::Debug, S> std::fmt::Debug for MemoMap<K, V, S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let inner = lock!(self.inner);
f.debug_struct("MemoMap")
.field("inner", &DebugMap(&inner))
.finish()
}
}
unsafe impl<K: Send, V: Send, S: Send> Send for MemoMap<K, V, S> {}
unsafe impl<K: Send + Sync, V: Send + Sync, S: Send> Sync for MemoMap<K, V, S> {}
impl<K: Clone, V: Clone, S: Clone> Clone for MemoMap<K, V, S> {
fn clone(&self) -> Self {
Self {
inner: Mutex::new(lock!(self.inner).clone()),
}
}
}
impl<K, V, S: Default> Default for MemoMap<K, V, S> {
fn default() -> Self {
MemoMap {
inner: Mutex::new(HashMap::default()),
}
}
}
impl<K, V> MemoMap<K, V, RandomState> {
pub fn new() -> MemoMap<K, V, RandomState> {
MemoMap {
inner: Mutex::default(),
}
}
}
impl<K, V, S> MemoMap<K, V, S> {
pub fn with_hasher(hash_builder: S) -> MemoMap<K, V, S> {
MemoMap {
inner: Mutex::new(HashMap::with_hasher(hash_builder)),
}
}
}
impl<K, V, S> MemoMap<K, V, S>
where
K: Eq + Hash,
S: BuildHasher,
{
pub fn insert(&self, key: K, value: V) -> bool {
let mut inner = lock!(self.inner);
match inner.entry(StableEntry::new(key, value)) {
Entry::Occupied(_) => false,
Entry::Vacant(vacant) => {
vacant.insert(());
true
}
}
}
pub fn replace(&mut self, key: K, value: V) {
let mut inner = lock!(self.inner);
if let Some((entry, _)) = inner.get_key_value(BorrowedKey::from_ref(&key)) {
let old_value = unsafe { std::mem::replace(&mut *entry.value_ptr(), value) };
drop(old_value);
} else {
inner.insert(StableEntry::new(key, value), ());
}
}
pub fn contains_key<Q>(&self, key: &Q) -> bool
where
Q: Hash + Eq + ?Sized,
K: Borrow<Q>,
{
lock!(self.inner).contains_key(BorrowedKey::from_ref(key))
}
pub fn get<Q>(&self, key: &Q) -> Option<&V>
where
Q: Hash + Eq + ?Sized,
K: Borrow<Q>,
{
let inner = lock!(self.inner);
let value = inner.get_key_value(BorrowedKey::from_ref(key))?.0.value();
Some(unsafe { transmute::<&V, &V>(value) })
}
#[allow(clippy::mutable_key_type)] pub fn get_mut<Q>(&mut self, key: &Q) -> Option<&mut V>
where
Q: Hash + Eq + ?Sized,
K: Borrow<Q>,
{
get_mut!(let map, self.inner);
let entry = map.get_key_value(BorrowedKey::from_ref(key))?.0;
Some(unsafe { &mut *entry.value_ptr() })
}
pub fn get_or_try_insert<Q, F, E>(&self, key: &Q, creator: F) -> Result<&V, E>
where
Q: Hash + Eq + ToOwned<Owned = K> + ?Sized,
K: Borrow<Q>,
F: FnOnce() -> Result<V, E>,
{
let mut inner = lock!(self.inner);
let key = BorrowedKey::from_ref(key);
if let Some((entry, _)) = inner.get_key_value(key) {
return Ok(unsafe { transmute::<&V, &V>(entry.value()) });
}
let entry = StableEntry::new(key.0.to_owned(), creator()?);
let value_ptr = match inner.entry(entry) {
Entry::Occupied(entry) => entry.key().value_ptr(),
Entry::Vacant(entry) => {
let value_ptr = entry.key().value_ptr();
entry.insert(());
value_ptr
}
};
Ok(unsafe { &*value_ptr })
}
pub fn get_or_insert_owned<F>(&self, key: K, creator: F) -> &V
where
F: FnOnce() -> V,
{
self.get_or_try_insert_owned(key, || Ok::<_, Infallible>(creator()))
.unwrap()
}
pub fn get_or_try_insert_owned<F, E>(&self, key: K, creator: F) -> Result<&V, E>
where
F: FnOnce() -> Result<V, E>,
{
let mut inner = lock!(self.inner);
if let Some((entry, _)) = inner.get_key_value(BorrowedKey::from_ref(&key)) {
return Ok(unsafe { transmute::<&V, &V>(entry.value()) });
}
let entry = StableEntry::new(key, creator()?);
let value_ptr = match inner.entry(entry) {
Entry::Occupied(entry) => entry.key().value_ptr(),
Entry::Vacant(entry) => {
let value_ptr = entry.key().value_ptr();
entry.insert(());
value_ptr
}
};
Ok(unsafe { &*value_ptr })
}
pub fn get_or_insert<Q, F>(&self, key: &Q, creator: F) -> &V
where
Q: Hash + Eq + ToOwned<Owned = K> + ?Sized,
K: Borrow<Q>,
F: FnOnce() -> V,
{
self.get_or_try_insert(key, || Ok::<_, Infallible>(creator()))
.unwrap()
}
pub fn remove<Q>(&mut self, key: &Q) -> Option<V>
where
Q: Hash + Eq + ?Sized,
K: Borrow<Q>,
{
lock!(self.inner)
.remove_entry(BorrowedKey::from_ref(key))
.map(|(entry, ())| entry.into_value())
}
pub fn clear(&mut self) {
lock!(self.inner).clear();
}
pub fn len(&self) -> usize {
lock!(self.inner).len()
}
pub fn is_empty(&self) -> bool {
lock!(self.inner).is_empty()
}
pub fn iter(&self) -> Iter<'_, K, V, S> {
let guard = lock!(self.inner);
let iter = guard.iter();
Iter {
iter: ManuallyDrop::new(unsafe {
transmute::<
hash_map::Iter<'_, StableEntry<K, V>, ()>,
hash_map::Iter<'_, StableEntry<K, V>, ()>,
>(iter)
}),
guard: ManuallyDrop::new(guard),
}
}
#[allow(clippy::mutable_key_type)] pub fn iter_mut(&mut self) -> IterMut<'_, K, V> {
get_mut!(let map, self.inner);
IterMut {
iter: unsafe {
transmute::<
hash_map::IterMut<'_, StableEntry<K, V>, ()>,
hash_map::IterMut<'_, StableEntry<K, V>, ()>,
>(map.iter_mut())
},
}
}
#[allow(clippy::mutable_key_type)] pub fn values_mut(&mut self) -> ValuesMut<'_, K, V> {
get_mut!(let map, self.inner);
ValuesMut {
iter: unsafe {
transmute::<
hash_map::IterMut<'_, StableEntry<K, V>, ()>,
hash_map::IterMut<'_, StableEntry<K, V>, ()>,
>(map.iter_mut())
},
}
}
pub fn keys(&self) -> Keys<'_, K, V, S> {
Keys { iter: self.iter() }
}
}
pub struct Iter<'a, K, V, S> {
iter: ManuallyDrop<hash_map::Iter<'a, StableEntry<K, V>, ()>>,
guard: ManuallyDrop<MutexGuard<'a, InnerMap<K, V, S>>>,
}
impl<K, V, S> Drop for Iter<'_, K, V, S> {
fn drop(&mut self) {
unsafe {
ManuallyDrop::drop(&mut self.iter);
ManuallyDrop::drop(&mut self.guard);
}
}
}
impl<'a, K, V, S> Iterator for Iter<'a, K, V, S> {
type Item = (&'a K, &'a V);
fn next(&mut self) -> Option<Self::Item> {
self.iter
.next()
.map(|(entry, ())| (entry.key(), entry.value()))
}
}
pub struct Keys<'a, K, V, S> {
iter: Iter<'a, K, V, S>,
}
impl<'a, K, V, S> Iterator for Keys<'a, K, V, S> {
type Item = &'a K;
fn next(&mut self) -> Option<Self::Item> {
self.iter.next().map(|(k, _)| k)
}
}
pub struct IterMut<'a, K, V> {
iter: hash_map::IterMut<'a, StableEntry<K, V>, ()>,
}
impl<'a, K, V> Iterator for IterMut<'a, K, V> {
type Item = (&'a K, &'a mut V);
fn next(&mut self) -> Option<Self::Item> {
self.iter.next().map(|(entry, ())| {
(entry.key(), unsafe { &mut *entry.value_ptr() })
})
}
}
pub struct ValuesMut<'a, K, V> {
iter: hash_map::IterMut<'a, StableEntry<K, V>, ()>,
}
impl<'a, K, V> Iterator for ValuesMut<'a, K, V> {
type Item = &'a mut V;
fn next(&mut self) -> Option<Self::Item> {
self.iter.next().map(|(entry, ())| {
unsafe { &mut *entry.value_ptr() }
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_insert() {
let memo = MemoMap::new();
assert!(memo.insert(23u32, Box::new(1u32)));
assert!(!memo.insert(23u32, Box::new(2u32)));
assert_eq!(memo.get(&23u32).cloned(), Some(Box::new(1)));
}
#[test]
fn test_iter() {
let memo = MemoMap::new();
memo.insert(1, "one");
memo.insert(2, "two");
memo.insert(3, "three");
let mut values = memo.iter().map(|(k, v)| (*k, *v)).collect::<Vec<_>>();
values.sort();
assert_eq!(values, vec![(1, "one"), (2, "two"), (3, "three")]);
}
#[test]
fn test_keys() {
let memo = MemoMap::new();
memo.insert(1, "one");
memo.insert(2, "two");
memo.insert(3, "three");
let mut values = memo.keys().copied().collect::<Vec<_>>();
values.sort();
assert_eq!(values, vec![1, 2, 3]);
}
#[test]
fn test_contains() {
let memo = MemoMap::new();
memo.insert(1, "one");
assert!(memo.contains_key(&1));
assert!(!memo.contains_key(&2));
}
#[test]
fn test_remove() {
let mut memo = MemoMap::new();
memo.insert(1, "one");
let value = memo.get(&1);
assert!(value.is_some());
let old_value = memo.remove(&1);
assert_eq!(old_value, Some("one"));
let value = memo.get(&1);
assert!(value.is_none());
}
#[test]
fn test_clear() {
let mut memo = MemoMap::new();
memo.insert(1, "one");
memo.insert(2, "two");
assert_eq!(memo.len(), 2);
assert!(!memo.is_empty());
memo.clear();
assert_eq!(memo.len(), 0);
assert!(memo.is_empty());
}
#[test]
fn test_ref_after_resize() {
let memo = MemoMap::new();
let mut refs = Vec::new();
let iterations = if cfg!(miri) { 100 } else { 10000 };
for key in 0..iterations {
refs.push((key, memo.get_or_insert(&key, || Box::new(key))));
}
for (key, val) in refs {
dbg!(key, val);
assert_eq!(memo.get(&key), Some(val));
}
}
#[test]
fn test_ref_after_resize_owned() {
let memo = MemoMap::new();
let mut refs = Vec::new();
let iterations = if cfg!(miri) { 100 } else { 10000 };
for key in 0..iterations {
refs.push((
key,
memo.get_or_insert_owned(key.to_string(), || Box::new(key)),
));
}
for (key, val) in refs {
dbg!(key, val);
assert_eq!(memo.get(&key.to_string()), Some(val));
}
}
#[test]
fn test_key_ref_after_resize() {
let memo = MemoMap::new();
memo.insert(0usize, 0usize);
let key = memo.keys().next().unwrap();
let iterations = if cfg!(miri) { 100 } else { 10000 };
for value in 1..iterations {
memo.insert(value, value);
}
assert_eq!(*key, 0);
}
#[test]
fn test_borrowed_key_lookup() {
let mut memo = MemoMap::new();
memo.insert("key".to_string(), 42);
assert!(memo.contains_key("key"));
assert_eq!(memo.get("key"), Some(&42));
*memo.get_mut("key").unwrap() = 43;
assert_eq!(memo.remove("key"), Some(43));
}
#[test]
fn test_replace() {
let mut memo = MemoMap::new();
memo.insert("foo", "bar");
memo.replace("foo", "bar2");
assert_eq!(memo.get("foo"), Some(&"bar2"));
}
#[test]
fn test_get_mut() {
let mut memo = MemoMap::new();
memo.insert("foo", "bar");
*memo.get_mut("foo").unwrap() = "bar2";
assert_eq!(memo.get("foo"), Some(&"bar2"));
}
#[test]
fn test_iter_mut() {
let mut memo = MemoMap::new();
memo.insert("foo", "bar");
for item in memo.iter_mut() {
*item.1 = "bar2";
}
assert_eq!(memo.get("foo"), Some(&"bar2"));
}
#[test]
fn test_values_mut() {
let mut memo = MemoMap::new();
memo.insert("foo", "bar");
for item in memo.values_mut() {
*item = "bar2";
}
assert_eq!(memo.get("foo"), Some(&"bar2"));
}
}