Skip to main content

moirai_sync/sync/
resource_pool.rs

1use std::collections::{hash_map::DefaultHasher, VecDeque};
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            static THREAD_SHARD_INDEX: std::cell::Cell<Option<usize>> = const { std::cell::Cell::new(None) };
91        }
92        THREAD_SHARD_INDEX.with(|cell| {
93            if let Some(idx) = cell.get() {
94                idx
95            } else {
96                let thread_id = std::thread::current().id();
97                let mut hasher = DefaultHasher::new();
98                thread_id.hash(&mut hasher);
99                let idx = (hasher.finish() as usize) % 4;
100                cell.set(Some(idx));
101                idx
102            }
103        })
104    }
105
106    /// Retrieve a resource of size >= `size` from the pool, or return `None`.
107    pub fn take_at_least(&self, size: u64) -> Option<T> {
108        let local_idx = Self::get_shard_index();
109        let start_bin = bin_index(size);
110
111        // Try local shard first
112        let local_shard = &self.shards[local_idx];
113
114        if local_shard.retained_count.load(Ordering::Acquire) > 0
115            && local_shard.retained_bytes.load(Ordering::Acquire) >= size
116        {
117            // 1. Search the start_bin for a buffer >= size (since start_bin contains elements of varying sizes)
118            {
119                let mut guard = local_shard.bins[start_bin].lock();
120                if let Some(pos) = guard.iter().rposition(|item| item.size() >= size) {
121                    let item = guard.remove(pos).expect("element exists at pos");
122                    let item_size = item.size();
123                    local_shard
124                        .retained_bytes
125                        .fetch_sub(item_size, Ordering::Release);
126                    local_shard.retained_count.fetch_sub(1, Ordering::Release);
127                    return Some(item);
128                }
129            }
130
131            // 2. Search larger bins (all elements in larger bins are guaranteed to be >= size)
132            for b in (start_bin + 1)..64 {
133                let mut guard = local_shard.bins[b].lock();
134                if let Some(item) = guard.pop_back() {
135                    let item_size = item.size();
136                    local_shard
137                        .retained_bytes
138                        .fetch_sub(item_size, Ordering::Release);
139                    local_shard.retained_count.fetch_sub(1, Ordering::Release);
140                    return Some(item);
141                }
142            }
143        }
144
145        // Steal from other shards using non-blocking try_lock
146        for i in 1..4 {
147            let other_idx = (local_idx + i) % 4;
148            let other_shard = &self.shards[other_idx];
149
150            // Fast path check: if the other shard does not have any items or does not have enough bytes, skip it.
151            if other_shard.retained_count.load(Ordering::Acquire) == 0
152                || other_shard.retained_bytes.load(Ordering::Acquire) < size
153            {
154                continue;
155            }
156
157            // 1. Search start_bin of other shard
158            if let Some(mut guard) = other_shard.bins[start_bin].try_lock() {
159                if let Some(pos) = guard.iter().rposition(|item| item.size() >= size) {
160                    let item = guard.remove(pos).expect("element exists at pos");
161                    let item_size = item.size();
162                    other_shard
163                        .retained_bytes
164                        .fetch_sub(item_size, Ordering::Release);
165                    other_shard.retained_count.fetch_sub(1, Ordering::Release);
166                    return Some(item);
167                }
168            }
169
170            // 2. Search larger bins of other shard
171            for b in (start_bin + 1)..64 {
172                if let Some(mut guard) = other_shard.bins[b].try_lock() {
173                    if let Some(item) = guard.pop_back() {
174                        let item_size = item.size();
175                        other_shard
176                            .retained_bytes
177                            .fetch_sub(item_size, Ordering::Release);
178                        other_shard.retained_count.fetch_sub(1, Ordering::Release);
179                        return Some(item);
180                    }
181                }
182            }
183        }
184
185        None
186    }
187
188    /// Recycle a resource back into the pool.
189    pub fn recycle(&self, item: T) {
190        let size = item.size();
191        if size > self.shard_max_bytes || self.shard_max_buffers == 0 {
192            return;
193        }
194
195        let local_idx = Self::get_shard_index();
196        let local_shard = &self.shards[local_idx];
197        let bin_idx = bin_index(size);
198
199        // The target-bin guard covers reservation through publication. `clear`
200        // acquires every bin guard before draining or resetting counters, so it
201        // cannot publish a zero-counter state between these two mutations.
202        let mut target_guard = local_shard.bins[bin_idx].lock();
203
204        // Reserve this item's count and bytes up front, before inserting, so the
205        // eviction decision below sees a total that already includes this item
206        // *and* every other concurrent recycler's in-flight contribution. The
207        // prior load-decide-insert sequence read the counters, decided no
208        // eviction was needed, then inserted — allowing N concurrent recyclers to
209        // each skip eviction and overshoot the shard cap by up to N-1 buffers
210        // (and exceed the byte budget). `fetch_add` returns the pre-add value, so
211        // `+ 1` / `+ size` is this shard's total with our reservation applied.
212        let mut current_count = local_shard.retained_count.fetch_add(1, Ordering::AcqRel) + 1;
213        let mut current_bytes = local_shard.retained_bytes.fetch_add(size, Ordering::AcqRel) + size;
214
215        // Evict oldest items (FIFO) until the shard — counting our reserved item —
216        // is within both limits, or no further eviction is possible. The local
217        // `current_*` counters are decremented per eviction (rather than
218        // re-loaded) so the loop terminates under sustained concurrent recycling
219        // instead of chasing a moving atomic snapshot; a single item always fits
220        // because `size <= shard_max_bytes` and `shard_max_buffers >= 1`.
221        let mut evicted = Vec::new();
222        while current_count > self.shard_max_buffers || current_bytes > self.shard_max_bytes {
223            let mut progress = false;
224            for b in 0..64 {
225                if b == bin_idx {
226                    if let Some(removed) = target_guard.pop_front() {
227                        let removed_size = removed.size();
228                        // Decrements remove already-inserted items, never our
229                        // reservation, so the net total keeps counting our item.
230                        local_shard.retained_count.fetch_sub(1, Ordering::Release);
231                        local_shard
232                            .retained_bytes
233                            .fetch_sub(removed_size, Ordering::Release);
234                        current_count -= 1;
235                        current_bytes = current_bytes.saturating_sub(removed_size);
236                        evicted.push(removed);
237                        progress = true;
238                        break;
239                    }
240                } else if let Some(mut guard) = local_shard.bins[b].try_lock() {
241                    if let Some(removed) = guard.pop_front() {
242                        let removed_size = removed.size();
243                        // Decrements remove already-inserted items, never our
244                        // reservation, so the net total keeps counting our item.
245                        local_shard.retained_count.fetch_sub(1, Ordering::Release);
246                        local_shard
247                            .retained_bytes
248                            .fetch_sub(removed_size, Ordering::Release);
249                        current_count -= 1;
250                        current_bytes = current_bytes.saturating_sub(removed_size);
251                        evicted.push(removed);
252                        progress = true;
253                        break;
254                    }
255                }
256            }
257            if !progress {
258                break;
259            }
260        }
261
262        #[cfg(test)]
263        self.test_hook.pause_after_reservation(local_idx, bin_idx);
264
265        // The counters already account for this item (reserved above); inserting
266        // it makes the bin contents consistent with the published totals.
267        target_guard.push_back(item);
268        drop(target_guard);
269        drop(evicted);
270    }
271
272    /// Clear all pooled resources.
273    ///
274    /// All bin guards remain held until the bins are drained and the counters
275    /// are reset. This makes the reset a linearization point: a concurrent
276    /// `recycle` or `take_at_least` either completes before the reset or starts
277    /// after it, and cannot publish a resource behind zero counters.
278    pub fn clear(&self) {
279        for (shard_idx, shard) in self.shards.iter().enumerate() {
280            #[cfg(not(test))]
281            let _ = shard_idx;
282            let mut guards: [Option<_>; 64] = std::array::from_fn(|_| None);
283            for (bin_idx, bin) in shard.bins.iter().enumerate() {
284                #[cfg(test)]
285                self.test_hook.announce_clear(shard_idx, bin_idx);
286                guards[bin_idx] = Some(bin.lock());
287            }
288
289            let mut evicted = Vec::new();
290            for guard in guards.iter_mut().flatten() {
291                evicted.extend(guard.drain(..));
292            }
293            shard.retained_bytes.store(0, Ordering::Release);
294            shard.retained_count.store(0, Ordering::Release);
295
296            drop(guards);
297            drop(evicted);
298        }
299    }
300
301    #[cfg(test)]
302    pub(crate) fn install_test_hook(
303        &self,
304        recycle_entered: std::sync::mpsc::SyncSender<()>,
305        clear_started: std::sync::mpsc::SyncSender<()>,
306        release: std::sync::Arc<std::sync::Barrier>,
307    ) -> test_support::HookGuard {
308        self.test_hook
309            .install(recycle_entered, clear_started, release)
310    }
311}
312
313#[cfg(test)]
314pub(crate) mod test_support {
315    use std::sync::{mpsc::SyncSender, Arc, Barrier, Mutex};
316
317    struct InterleavingHook {
318        recycle_entered: SyncSender<()>,
319        clear_started: SyncSender<()>,
320        release: Arc<Barrier>,
321        target: Option<(usize, usize)>,
322        clear_announced: bool,
323    }
324
325    pub(crate) struct Hook {
326        state: Arc<Mutex<Option<InterleavingHook>>>,
327    }
328
329    impl Hook {
330        pub(crate) fn new() -> Self {
331            Self {
332                state: Arc::new(Mutex::new(None)),
333            }
334        }
335
336        pub(crate) fn install(
337            &self,
338            recycle_entered: SyncSender<()>,
339            clear_started: SyncSender<()>,
340            release: Arc<Barrier>,
341        ) -> HookGuard {
342            let mut hook = self
343                .state
344                .lock()
345                .expect("invariant: test hook mutex poisoned");
346            assert!(
347                hook.is_none(),
348                "invariant: only one interleaving hook is active"
349            );
350            *hook = Some(InterleavingHook {
351                recycle_entered,
352                clear_started,
353                release,
354                target: None,
355                clear_announced: false,
356            });
357            HookGuard {
358                state: Arc::clone(&self.state),
359            }
360        }
361
362        pub(crate) fn pause_after_reservation(&self, shard_idx: usize, bin_idx: usize) {
363            let (entered, release) = {
364                let mut hook = self
365                    .state
366                    .lock()
367                    .expect("invariant: test hook mutex poisoned");
368                let Some(hook) = hook.as_mut() else {
369                    return;
370                };
371                assert!(
372                    hook.target.replace((shard_idx, bin_idx)).is_none(),
373                    "invariant: only one recycle interleaving is active"
374                );
375                (hook.recycle_entered.clone(), Arc::clone(&hook.release))
376            };
377
378            entered
379                .send(())
380                .expect("invariant: interleaving test receiver remains active");
381            release.wait();
382        }
383
384        pub(crate) fn announce_clear(&self, shard_idx: usize, bin_idx: usize) {
385            let started = {
386                let mut hook = self
387                    .state
388                    .lock()
389                    .expect("invariant: test hook mutex poisoned");
390                let Some(hook) = hook.as_mut() else {
391                    return;
392                };
393                if hook.target == Some((shard_idx, bin_idx)) && !hook.clear_announced {
394                    hook.clear_announced = true;
395                    Some(hook.clear_started.clone())
396                } else {
397                    None
398                }
399            };
400
401            if let Some(started) = started {
402                started
403                    .send(())
404                    .expect("invariant: interleaving test receiver remains active");
405            }
406        }
407    }
408
409    pub(crate) struct HookGuard {
410        state: Arc<Mutex<Option<InterleavingHook>>>,
411    }
412
413    impl Drop for HookGuard {
414        fn drop(&mut self) {
415            self.state
416                .lock()
417                .expect("invariant: test hook mutex poisoned")
418                .take();
419        }
420    }
421}