moirai_sync/sync/
concurrent_hash_map.rs1use std::collections::HashMap;
2use std::collections::hash_map::RandomState;
3use std::fmt;
4use std::hash::{BuildHasher, Hash, Hasher};
5use std::sync::RwLock;
6
7const DEFAULT_CONCURRENT_MAP_SEGMENTS: usize = 16;
9
10#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub struct SegmentPoisoned {
18 pub segment: usize,
20}
21
22impl fmt::Display for SegmentPoisoned {
23 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
24 write!(
25 f,
26 "concurrent map segment {} poisoned by a panicked writer",
27 self.segment
28 )
29 }
30}
31
32impl std::error::Error for SegmentPoisoned {}
33
34pub struct ConcurrentHashMap<K, V, S = RandomState> {
37 pub(crate) segments: Vec<RwLock<HashMap<K, V, S>>>,
38 pub(crate) hasher: S,
39}
40
41impl<K, V, S> fmt::Debug for ConcurrentHashMap<K, V, S> {
42 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
43 f.debug_struct("ConcurrentHashMap")
44 .field("segments_count", &self.segments.len())
45 .finish_non_exhaustive()
46 }
47}
48
49impl<K: Hash + Eq, V> ConcurrentHashMap<K, V> {
50 pub fn new() -> Self {
52 Self::with_segments(DEFAULT_CONCURRENT_MAP_SEGMENTS)
53 }
54
55 pub fn with_segments(num_segments: usize) -> Self {
57 let num_segments = num_segments.next_power_of_two();
58
59 let segments = (0..num_segments)
60 .map(|_| RwLock::new(HashMap::new()))
61 .collect();
62
63 Self {
64 segments,
65 hasher: RandomState::new(),
66 }
67 }
68}
69
70impl<K: Hash + Eq, V> Default for ConcurrentHashMap<K, V> {
71 fn default() -> Self {
72 Self::new()
73 }
74}
75
76impl<K: Hash + Eq, V, S: BuildHasher> ConcurrentHashMap<K, V, S> {
77 pub(crate) fn segment_index(&self, key: &K) -> usize {
79 let mut hasher = self.hasher.build_hasher();
80 key.hash(&mut hasher);
81 let hash = hasher.finish();
82 (hash as usize) & (self.segments.len() - 1)
84 }
85
86 pub fn insert(&self, key: K, value: V) -> Result<Option<V>, SegmentPoisoned> {
95 let idx = self.segment_index(&key);
96 Ok(self.segments[idx]
97 .write()
98 .map_err(|_| SegmentPoisoned { segment: idx })?
99 .insert(key, value))
100 }
101
102 pub fn get_or_insert_with<F>(&self, key: K, default: F) -> Result<V, SegmentPoisoned>
113 where
114 F: FnOnce() -> V,
115 V: Clone,
116 {
117 let idx = self.segment_index(&key);
118 {
120 let shard = self.segments[idx]
121 .read()
122 .map_err(|_| SegmentPoisoned { segment: idx })?;
123 if let Some(value) = shard.get(&key) {
124 return Ok(value.clone());
125 }
126 }
127
128 let mut shard = self.segments[idx]
130 .write()
131 .map_err(|_| SegmentPoisoned { segment: idx })?;
132 if let Some(value) = shard.get(&key) {
133 Ok(value.clone())
134 } else {
135 let value = default();
136 shard.insert(key, value.clone());
137 Ok(value)
138 }
139 }
140
141 pub fn get(&self, key: &K) -> Result<Option<V>, SegmentPoisoned>
150 where
151 V: Clone,
152 {
153 let idx = self.segment_index(key);
154 Ok(self.segments[idx]
155 .read()
156 .map_err(|_| SegmentPoisoned { segment: idx })?
157 .get(key)
158 .cloned())
159 }
160
161 pub fn remove(&self, key: &K) -> Result<Option<V>, SegmentPoisoned> {
170 let idx = self.segment_index(key);
171 Ok(self.segments[idx]
172 .write()
173 .map_err(|_| SegmentPoisoned { segment: idx })?
174 .remove(key))
175 }
176
177 pub fn contains_key(&self, key: &K) -> Result<bool, SegmentPoisoned> {
186 let idx = self.segment_index(key);
187 Ok(self.segments[idx]
188 .read()
189 .map_err(|_| SegmentPoisoned { segment: idx })?
190 .contains_key(key))
191 }
192}