Skip to main content

moirai_sync/sync/
resource_pool.rs

1use std::collections::{VecDeque, hash_map::DefaultHasher};
2use std::hash::{Hash, Hasher};
3use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
4
5use crate::sync::spin_lock::SpinLock;
6
7/// A resource with a queryable byte size or capacity.
8pub trait SizeBounded {
9    /// The size or capacity of the resource in bytes.
10    fn size(&self) -> u64;
11}
12
13/// Helper to map an arbitrary size to its power-of-two bin index.
14#[inline]
15fn bin_index(size: u64) -> usize {
16    if size <= 1 {
17        0
18    } else {
19        // Subtract 1 so that exact powers of two fall into the exact bin,
20        // e.g. 1024 -> 10, 1025 -> 11.
21        64 - (size - 1).leading_zeros() as usize
22    }
23}
24
25struct Shard<T> {
26    // 64 bins, representing power-of-two size classes (2^0 to 2^63).
27    // SpinLock is cache-line aligned, preventing false sharing.
28    bins: [SpinLock<VecDeque<T>>; 64],
29    retained_bytes: AtomicU64,
30    retained_count: AtomicUsize,
31}
32
33impl<T> Shard<T> {
34    fn new() -> Self {
35        let mut bins_vec = Vec::with_capacity(64);
36        for _ in 0..64 {
37            bins_vec.push(SpinLock::new(VecDeque::new()));
38        }
39        let bins: [SpinLock<VecDeque<T>>; 64] = bins_vec
40            .try_into()
41            .unwrap_or_else(|_| panic!("invariant: failed to convert vector of 64 bins"));
42
43        Self {
44            bins,
45            retained_bytes: AtomicU64::new(0),
46            retained_count: AtomicUsize::new(0),
47        }
48    }
49}
50
51/// A sharded, binned resource pool designed for high-concurrency reuse of transient allocations.
52///
53/// Resources are partitioned by thread affinity across 4 shards to minimize lock contention,
54/// and internally binned into 64 power-of-two size classes. Pop operations use a non-blocking
55/// stealing fallback across shards.
56pub struct ShardedResourcePool<T> {
57    shards: [Shard<T>; 4],
58    shard_max_buffers: usize,
59    shard_max_bytes: u64,
60    #[cfg(test)]
61    test_hook: test_support::Hook,
62}
63
64impl<T> std::fmt::Debug for ShardedResourcePool<T> {
65    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
66        f.debug_struct("ShardedResourcePool")
67            .field("shard_max_buffers", &self.shard_max_buffers)
68            .field("shard_max_bytes", &self.shard_max_bytes)
69            .finish_non_exhaustive()
70    }
71}
72
73impl<T: SizeBounded> ShardedResourcePool<T> {
74    /// Construct a new pool with the given capacity limits.
75    #[must_use]
76    pub fn new(max_buffers: usize, max_bytes: u64) -> Self {
77        Self {
78            shards: [Shard::new(), Shard::new(), Shard::new(), Shard::new()],
79            shard_max_buffers: (max_buffers / 4).max(1),
80            shard_max_bytes: max_bytes / 4,
81            #[cfg(test)]
82            test_hook: test_support::Hook::new(),
83        }
84    }
85
86    /// Retrieve the thread-local shard index.
87    #[inline]
88    fn get_shard_index() -> usize {
89        thread_local! {
90            // clippy 1.97.0 false positive: initialiser is already
91            // `const { Cell::new(None) }`. Retire when toolchain advances
92            // past the regression (ATLAS-MNEMOSYNE-CI-1).
93            #[allow(clippy::missing_const_for_thread_local)]
94            static THREAD_SHARD_INDEX: std::cell::Cell<Option<usize>> = const { std::cell::Cell::new(None) };
95        }
96        THREAD_SHARD_INDEX.with(|cell| {
97            if let Some(idx) = cell.get() {
98                idx
99            } else {
100                let thread_id = std::thread::current().id();
101                let mut hasher = DefaultHasher::new();
102                thread_id.hash(&mut hasher);
103                let idx = (hasher.finish() as usize) % 4;
104                cell.set(Some(idx));
105                idx
106            }
107        })
108    }
109
110    /// Retrieve a resource of size >= `size` from the pool, or return `None`.
111    pub fn take_at_least(&self, size: u64) -> Option<T> {
112        let local_idx = Self::get_shard_index();
113        let start_bin = bin_index(size);
114
115        // Try local shard first
116        let local_shard = &self.shards[local_idx];
117
118        if local_shard.retained_count.load(Ordering::Acquire) > 0
119            && local_shard.retained_bytes.load(Ordering::Acquire) >= size
120        {
121            // 1. Search the start_bin for a buffer >= size (since start_bin contains elements of varying sizes)
122            {
123                let mut guard = local_shard.bins[start_bin].lock();
124                if let Some(pos) = guard.iter().rposition(|item| item.size() >= size) {
125                    let item = guard.remove(pos).expect("element exists at pos");
126                    let item_size = item.size();
127                    local_shard
128                        .retained_bytes
129                        .fetch_sub(item_size, Ordering::Release);
130                    local_shard.retained_count.fetch_sub(1, Ordering::Release);
131                    return Some(item);
132                }
133            }
134
135            // 2. Search larger bins (all elements in larger bins are guaranteed to be >= size)
136            for b in (start_bin + 1)..64 {
137                let mut guard = local_shard.bins[b].lock();
138                if let Some(item) = guard.pop_back() {
139                    let item_size = item.size();
140                    local_shard
141                        .retained_bytes
142                        .fetch_sub(item_size, Ordering::Release);
143                    local_shard.retained_count.fetch_sub(1, Ordering::Release);
144                    return Some(item);
145                }
146            }
147        }
148
149        // Steal from other shards using non-blocking try_lock
150        for i in 1..4 {
151            let other_idx = (local_idx + i) % 4;
152            let other_shard = &self.shards[other_idx];
153
154            // Fast path check: if the other shard does not have any items or does not have enough bytes, skip it.
155            if other_shard.retained_count.load(Ordering::Acquire) == 0
156                || other_shard.retained_bytes.load(Ordering::Acquire) < size
157            {
158                continue;
159            }
160
161            // 1. Search start_bin of other shard
162            if let Some(mut guard) = other_shard.bins[start_bin].try_lock()
163                && let Some(pos) = guard.iter().rposition(|item| item.size() >= size)
164            {
165                let item = guard.remove(pos).expect("element exists at pos");
166                let item_size = item.size();
167                other_shard
168                    .retained_bytes
169                    .fetch_sub(item_size, Ordering::Release);
170                other_shard.retained_count.fetch_sub(1, Ordering::Release);
171                return Some(item);
172            }
173
174            // 2. Search larger bins of other shard
175            for b in (start_bin + 1)..64 {
176                if let Some(mut guard) = other_shard.bins[b].try_lock()
177                    && let Some(item) = guard.pop_back()
178                {
179                    let item_size = item.size();
180                    other_shard
181                        .retained_bytes
182                        .fetch_sub(item_size, Ordering::Release);
183                    other_shard.retained_count.fetch_sub(1, Ordering::Release);
184                    return Some(item);
185                }
186            }
187        }
188
189        None
190    }
191
192    /// Recycle a resource back into the pool.
193    pub fn recycle(&self, item: T) {
194        let size = item.size();
195        if size > self.shard_max_bytes || self.shard_max_buffers == 0 {
196            return;
197        }
198
199        let local_idx = Self::get_shard_index();
200        let local_shard = &self.shards[local_idx];
201        let bin_idx = bin_index(size);
202
203        // The target-bin guard covers reservation through publication. `clear`
204        // acquires every bin guard before draining or resetting counters, so it
205        // cannot publish a zero-counter state between these two mutations.
206        let mut target_guard = local_shard.bins[bin_idx].lock();
207
208        // Reserve this item's count and bytes up front, before inserting, so the
209        // eviction decision below sees a total that already includes this item
210        // *and* every other concurrent recycler's in-flight contribution. The
211        // prior load-decide-insert sequence read the counters, decided no
212        // eviction was needed, then inserted — allowing N concurrent recyclers to
213        // each skip eviction and overshoot the shard cap by up to N-1 buffers
214        // (and exceed the byte budget). `fetch_add` returns the pre-add value, so
215        // `+ 1` / `+ size` is this shard's total with our reservation applied.
216        let mut current_count = local_shard.retained_count.fetch_add(1, Ordering::AcqRel) + 1;
217        let mut current_bytes = local_shard.retained_bytes.fetch_add(size, Ordering::AcqRel) + size;
218
219        // Evict oldest items (FIFO) until the shard — counting our reserved item —
220        // is within both limits, or no further eviction is possible. The local
221        // `current_*` counters are decremented per eviction (rather than
222        // re-loaded) so the loop terminates under sustained concurrent recycling
223        // instead of chasing a moving atomic snapshot; a single item always fits
224        // because `size <= shard_max_bytes` and `shard_max_buffers >= 1`.
225        let mut evicted = Vec::new();
226        while current_count > self.shard_max_buffers || current_bytes > self.shard_max_bytes {
227            let mut progress = false;
228            for b in 0..64 {
229                if b == bin_idx {
230                    if let Some(removed) = target_guard.pop_front() {
231                        let removed_size = removed.size();
232                        // Decrements remove already-inserted items, never our
233                        // reservation, so the net total keeps counting our item.
234                        local_shard.retained_count.fetch_sub(1, Ordering::Release);
235                        local_shard
236                            .retained_bytes
237                            .fetch_sub(removed_size, Ordering::Release);
238                        current_count -= 1;
239                        current_bytes = current_bytes.saturating_sub(removed_size);
240                        evicted.push(removed);
241                        progress = true;
242                        break;
243                    }
244                } else if let Some(mut guard) = local_shard.bins[b].try_lock()
245                    && let Some(removed) = guard.pop_front()
246                {
247                    let removed_size = removed.size();
248                    // Decrements remove already-inserted items, never our
249                    // reservation, so the net total keeps counting our item.
250                    local_shard.retained_count.fetch_sub(1, Ordering::Release);
251                    local_shard
252                        .retained_bytes
253                        .fetch_sub(removed_size, Ordering::Release);
254                    current_count -= 1;
255                    current_bytes = current_bytes.saturating_sub(removed_size);
256                    evicted.push(removed);
257                    progress = true;
258                    break;
259                }
260            }
261            if !progress {
262                break;
263            }
264        }
265
266        #[cfg(test)]
267        self.test_hook.pause_after_reservation(local_idx, bin_idx);
268
269        // The counters already account for this item (reserved above); inserting
270        // it makes the bin contents consistent with the published totals.
271        target_guard.push_back(item);
272        drop(target_guard);
273        drop(evicted);
274    }
275
276    /// Clear all pooled resources.
277    ///
278    /// All bin guards remain held until the bins are drained and the counters
279    /// are reset. This makes the reset a linearization point: a concurrent
280    /// `recycle` or `take_at_least` either completes before the reset or starts
281    /// after it, and cannot publish a resource behind zero counters.
282    pub fn clear(&self) {
283        for (shard_idx, shard) in self.shards.iter().enumerate() {
284            #[cfg(not(test))]
285            let _ = shard_idx;
286            let mut guards: [Option<_>; 64] = std::array::from_fn(|_| None);
287            for (bin_idx, bin) in shard.bins.iter().enumerate() {
288                #[cfg(test)]
289                self.test_hook.announce_clear(shard_idx, bin_idx);
290                guards[bin_idx] = Some(bin.lock());
291            }
292
293            let mut evicted = Vec::new();
294            for guard in guards.iter_mut().flatten() {
295                evicted.extend(guard.drain(..));
296            }
297            shard.retained_bytes.store(0, Ordering::Release);
298            shard.retained_count.store(0, Ordering::Release);
299
300            drop(guards);
301            drop(evicted);
302        }
303    }
304
305    #[cfg(test)]
306    pub(crate) fn install_test_hook(
307        &self,
308        recycle_entered: std::sync::mpsc::SyncSender<()>,
309        clear_started: std::sync::mpsc::SyncSender<()>,
310        release: std::sync::Arc<std::sync::Barrier>,
311    ) -> test_support::HookGuard {
312        self.test_hook
313            .install(recycle_entered, clear_started, release)
314    }
315}
316
317#[cfg(test)]
318pub(crate) mod test_support {
319    use std::sync::{Arc, Barrier, Mutex, mpsc::SyncSender};
320
321    struct InterleavingHook {
322        recycle_entered: SyncSender<()>,
323        clear_started: SyncSender<()>,
324        release: Arc<Barrier>,
325        target: Option<(usize, usize)>,
326        clear_announced: bool,
327    }
328
329    pub(crate) struct Hook {
330        state: Arc<Mutex<Option<InterleavingHook>>>,
331    }
332
333    impl Hook {
334        pub(crate) fn new() -> Self {
335            Self {
336                state: Arc::new(Mutex::new(None)),
337            }
338        }
339
340        pub(crate) fn install(
341            &self,
342            recycle_entered: SyncSender<()>,
343            clear_started: SyncSender<()>,
344            release: Arc<Barrier>,
345        ) -> HookGuard {
346            let mut hook = self
347                .state
348                .lock()
349                .expect("invariant: test hook mutex poisoned");
350            assert!(
351                hook.is_none(),
352                "invariant: only one interleaving hook is active"
353            );
354            *hook = Some(InterleavingHook {
355                recycle_entered,
356                clear_started,
357                release,
358                target: None,
359                clear_announced: false,
360            });
361            HookGuard {
362                state: Arc::clone(&self.state),
363            }
364        }
365
366        pub(crate) fn pause_after_reservation(&self, shard_idx: usize, bin_idx: usize) {
367            let (entered, release) = {
368                let mut hook = self
369                    .state
370                    .lock()
371                    .expect("invariant: test hook mutex poisoned");
372                let Some(hook) = hook.as_mut() else {
373                    return;
374                };
375                assert!(
376                    hook.target.replace((shard_idx, bin_idx)).is_none(),
377                    "invariant: only one recycle interleaving is active"
378                );
379                (hook.recycle_entered.clone(), Arc::clone(&hook.release))
380            };
381
382            entered
383                .send(())
384                .expect("invariant: interleaving test receiver remains active");
385            release.wait();
386        }
387
388        pub(crate) fn announce_clear(&self, shard_idx: usize, bin_idx: usize) {
389            let started = {
390                let mut hook = self
391                    .state
392                    .lock()
393                    .expect("invariant: test hook mutex poisoned");
394                let Some(hook) = hook.as_mut() else {
395                    return;
396                };
397                if hook.target == Some((shard_idx, bin_idx)) && !hook.clear_announced {
398                    hook.clear_announced = true;
399                    Some(hook.clear_started.clone())
400                } else {
401                    None
402                }
403            };
404
405            if let Some(started) = started {
406                started
407                    .send(())
408                    .expect("invariant: interleaving test receiver remains active");
409            }
410        }
411    }
412
413    pub(crate) struct HookGuard {
414        state: Arc<Mutex<Option<InterleavingHook>>>,
415    }
416
417    impl Drop for HookGuard {
418        fn drop(&mut self) {
419            self.state
420                .lock()
421                .expect("invariant: test hook mutex poisoned")
422                .take();
423        }
424    }
425}