moirai_sync/sync/
resource_pool.rs1use 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
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 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 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 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 {
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 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 for i in 1..4 {
147 let other_idx = (local_idx + i) % 4;
148 let other_shard = &self.shards[other_idx];
149
150 if other_shard.retained_count.load(Ordering::Acquire) == 0
152 || other_shard.retained_bytes.load(Ordering::Acquire) < size
153 {
154 continue;
155 }
156
157 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 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 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 let mut target_guard = local_shard.bins[bin_idx].lock();
203
204 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 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 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 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 target_guard.push_back(item);
268 drop(target_guard);
269 drop(evicted);
270 }
271
272 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}