use std::collections::HashMap;
use std::collections::hash_map::RandomState;
use std::fmt;
use std::hash::{BuildHasher, Hash, Hasher};
use std::sync::RwLock;
const DEFAULT_CONCURRENT_MAP_SEGMENTS: usize = 16;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SegmentPoisoned {
pub segment: usize,
}
impl fmt::Display for SegmentPoisoned {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"concurrent map segment {} poisoned by a panicked writer",
self.segment
)
}
}
impl std::error::Error for SegmentPoisoned {}
pub struct ConcurrentHashMap<K, V, S = RandomState> {
pub(crate) segments: Vec<RwLock<HashMap<K, V, S>>>,
pub(crate) hasher: S,
}
impl<K, V, S> fmt::Debug for ConcurrentHashMap<K, V, S> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ConcurrentHashMap")
.field("segments_count", &self.segments.len())
.finish_non_exhaustive()
}
}
impl<K: Hash + Eq, V> ConcurrentHashMap<K, V> {
pub fn new() -> Self {
Self::with_segments(DEFAULT_CONCURRENT_MAP_SEGMENTS)
}
pub fn with_segments(num_segments: usize) -> Self {
let num_segments = num_segments.next_power_of_two();
let segments = (0..num_segments)
.map(|_| RwLock::new(HashMap::new()))
.collect();
Self {
segments,
hasher: RandomState::new(),
}
}
}
impl<K: Hash + Eq, V> Default for ConcurrentHashMap<K, V> {
fn default() -> Self {
Self::new()
}
}
impl<K: Hash + Eq, V, S: BuildHasher> ConcurrentHashMap<K, V, S> {
pub(crate) fn segment_index(&self, key: &K) -> usize {
let mut hasher = self.hasher.build_hasher();
key.hash(&mut hasher);
let hash = hasher.finish();
(hash as usize) & (self.segments.len() - 1)
}
pub fn insert(&self, key: K, value: V) -> Result<Option<V>, SegmentPoisoned> {
let idx = self.segment_index(&key);
Ok(self.segments[idx]
.write()
.map_err(|_| SegmentPoisoned { segment: idx })?
.insert(key, value))
}
pub fn get_or_insert_with<F>(&self, key: K, default: F) -> Result<V, SegmentPoisoned>
where
F: FnOnce() -> V,
V: Clone,
{
let idx = self.segment_index(&key);
{
let shard = self.segments[idx]
.read()
.map_err(|_| SegmentPoisoned { segment: idx })?;
if let Some(value) = shard.get(&key) {
return Ok(value.clone());
}
}
let mut shard = self.segments[idx]
.write()
.map_err(|_| SegmentPoisoned { segment: idx })?;
if let Some(value) = shard.get(&key) {
Ok(value.clone())
} else {
let value = default();
shard.insert(key, value.clone());
Ok(value)
}
}
pub fn get(&self, key: &K) -> Result<Option<V>, SegmentPoisoned>
where
V: Clone,
{
let idx = self.segment_index(key);
Ok(self.segments[idx]
.read()
.map_err(|_| SegmentPoisoned { segment: idx })?
.get(key)
.cloned())
}
pub fn remove(&self, key: &K) -> Result<Option<V>, SegmentPoisoned> {
let idx = self.segment_index(key);
Ok(self.segments[idx]
.write()
.map_err(|_| SegmentPoisoned { segment: idx })?
.remove(key))
}
pub fn contains_key(&self, key: &K) -> Result<bool, SegmentPoisoned> {
let idx = self.segment_index(key);
Ok(self.segments[idx]
.read()
.map_err(|_| SegmentPoisoned { segment: idx })?
.contains_key(key))
}
}