Skip to main content

moirai_sync/sync/
concurrent_hash_map.rs

1use std::collections::HashMap;
2use std::collections::hash_map::RandomState;
3use std::fmt;
4use std::hash::{BuildHasher, Hash, Hasher};
5use std::sync::RwLock;
6
7/// Default concurrent map segment count for optimal performance
8const DEFAULT_CONCURRENT_MAP_SEGMENTS: usize = 16;
9
10/// Error returned when a segment's `RwLock` was poisoned by a panicked writer.
11///
12/// Carries the index of the poisoned segment so diagnostics can name the
13/// offending shard. This is a genuine contract failure (a writer panicked
14/// while holding the segment lock), so it is surfaced as a typed error
15/// rather than recovered silently.
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub struct SegmentPoisoned {
18    /// Index of the poisoned segment.
19    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
34/// Concurrent hash map with segment-based locking for scalability.
35/// This provides better scalability than a single mutex-protected HashMap.
36pub 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    /// Create a new concurrent hash map with default hasher.
51    pub fn new() -> Self {
52        Self::with_segments(DEFAULT_CONCURRENT_MAP_SEGMENTS)
53    }
54
55    /// Create with a specific number of segments (must be power of 2).
56    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    /// Get the segment index for a key.
78    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        // Use bitmask for even distribution across power-of-2 segments
83        (hash as usize) & (self.segments.len() - 1)
84    }
85
86    /// Insert a key-value pair.
87    ///
88    /// Returns the previous value if the key existed, or None if it was a new key.
89    ///
90    /// # Errors
91    ///
92    /// Returns [`SegmentPoisoned`] if the segment lock was poisoned by a
93    /// panicked writer.
94    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    /// Get a value by key, or insert it if it is not present.
103    ///
104    /// Executes atomically under the segment's write lock, ensuring that the
105    /// `default` closure runs exactly once on a cache miss and no concurrent insert
106    /// can overwrite it.
107    ///
108    /// # Errors
109    ///
110    /// Returns [`SegmentPoisoned`] if the segment lock was poisoned by a
111    /// panicked writer.
112    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        // Phase 1: Fast-path with read lock
119        {
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        // Phase 2: Slow-path with write lock
129        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    /// Get a value by key.
142    ///
143    /// Returns the cloned value if found, or None if not found.
144    ///
145    /// # Errors
146    ///
147    /// Returns [`SegmentPoisoned`] if the segment lock was poisoned by a
148    /// panicked writer.
149    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    /// Remove a key-value pair.
162    ///
163    /// Returns the removed value if the key existed, or None if it didn't exist.
164    ///
165    /// # Errors
166    ///
167    /// Returns [`SegmentPoisoned`] if the segment lock was poisoned by a
168    /// panicked writer.
169    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    /// Check if a key exists.
178    ///
179    /// Returns true if the key exists, false otherwise.
180    ///
181    /// # Errors
182    ///
183    /// Returns [`SegmentPoisoned`] if the segment lock was poisoned by a
184    /// panicked writer.
185    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}