moirai_sync/sync/
resource_pool.rs1use 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
7pub trait SizeBounded {
9 fn size(&self) -> u64;
11}
12
13#[inline]
15fn bin_index(size: u64) -> usize {
16 if size <= 1 {
17 0
18 } else {
19 64 - (size - 1).leading_zeros() as usize
22 }
23}
24
25struct Shard<T> {
26 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
51pub 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 #[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 #[inline]
88 fn get_shard_index() -> usize {
89 thread_local! {
90 #[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 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 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 {
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 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 for i in 1..4 {
151 let other_idx = (local_idx + i) % 4;
152 let other_shard = &self.shards[other_idx];
153
154 if other_shard.retained_count.load(Ordering::Acquire) == 0
156 || other_shard.retained_bytes.load(Ordering::Acquire) < size
157 {
158 continue;
159 }
160
161 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 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 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 let mut target_guard = local_shard.bins[bin_idx].lock();
207
208 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 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 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 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 target_guard.push_back(item);
272 drop(target_guard);
273 drop(evicted);
274 }
275
276 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}