binius_compute/
buffer_pool.rs1use std::{
17 collections::HashMap,
18 fmt,
19 mem::{self, MaybeUninit},
20 ops::{Deref, DerefMut},
21 sync::Mutex,
22};
23
24const BUFFER_ALIGN: usize = 64;
31
32#[repr(align(64))]
38struct AlignedChunk(#[allow(dead_code)] [u8; BUFFER_ALIGN]);
39
40#[derive(Default)]
48pub struct BufferPool {
49 free_list: Mutex<HashMap<usize, Vec<Vec<AlignedChunk>>>>,
50}
51
52impl BufferPool {
53 pub fn new() -> Self {
55 Self::default()
56 }
57
58 pub fn alloc_vec<T>(&self, capacity: usize) -> PoolVec<'_, T> {
69 PoolVec {
70 pool: self,
71 data: self.alloc_data(capacity),
72 }
73 }
74
75 fn alloc_data<T>(&self, capacity: usize) -> Vec<T> {
76 const {
77 assert!(
78 mem::align_of::<T>() <= BUFFER_ALIGN,
79 "element alignment exceeds the pool's buffer alignment"
80 );
81 assert!(
82 mem::size_of::<T>().is_power_of_two() || mem::size_of::<T>() == 0,
83 "element size must be a power of two for byte-size-keyed reuse"
84 );
85 }
86
87 if mem::size_of::<T>() == 0 || capacity == 0 {
90 return Vec::with_capacity(capacity);
91 }
92
93 let byte_len = (capacity * mem::size_of::<T>())
94 .next_power_of_two()
95 .max(BUFFER_ALIGN);
96 let n_chunks = byte_len / BUFFER_ALIGN;
97 let elem_cap = byte_len / mem::size_of::<T>();
98
99 let reused = self
100 .free_list
101 .lock()
102 .expect("free list mutex poisoned")
103 .get_mut(&n_chunks)
104 .and_then(Vec::pop);
105
106 let mut block = reused.unwrap_or_else(|| Vec::with_capacity(n_chunks));
107 let ptr = block.as_mut_ptr().cast::<T>();
108 mem::forget(block);
111 unsafe { Vec::from_raw_parts(ptr, 0, elem_cap) }
117 }
118
119 fn reclaim(&self, n_chunks: usize, block: Vec<AlignedChunk>) {
120 self.free_list
121 .lock()
122 .expect("free list mutex poisoned")
123 .entry(n_chunks)
124 .or_default()
125 .push(block);
126 }
127}
128
129impl fmt::Debug for BufferPool {
130 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
131 f.debug_struct("BufferPool").finish_non_exhaustive()
132 }
133}
134
135pub struct PoolVec<'alloc, T> {
141 pool: &'alloc BufferPool,
142 data: Vec<T>,
143}
144
145impl<T> PoolVec<'_, T> {
146 pub const fn capacity(&self) -> usize {
148 self.data.capacity()
149 }
150
151 pub fn push(&mut self, value: T) {
153 self.data.push(value);
154 }
155
156 pub fn clear(&mut self) {
158 self.data.clear();
159 }
160
161 pub fn truncate(&mut self, len: usize) {
165 self.data.truncate(len);
166 }
167
168 pub fn spare_capacity_mut(&mut self) -> &mut [MaybeUninit<T>] {
173 self.data.spare_capacity_mut()
174 }
175
176 pub unsafe fn set_len(&mut self, new_len: usize) {
183 unsafe { self.data.set_len(new_len) }
184 }
185}
186
187impl<T: Clone> PoolVec<'_, T> {
188 pub fn extend_from_slice(&mut self, other: &[T]) {
190 self.data.extend_from_slice(other);
191 }
192
193 pub fn resize(&mut self, new_len: usize, value: T) {
195 self.data.resize(new_len, value);
196 }
197}
198
199impl<T: Clone> Clone for PoolVec<'_, T> {
200 fn clone(&self) -> Self {
205 let mut cloned = self.pool.alloc_vec::<T>(self.data.len());
206 cloned.extend_from_slice(&self.data);
207 cloned
208 }
209}
210
211impl<T> Drop for PoolVec<'_, T> {
212 fn drop(&mut self) {
213 let mut data = mem::take(&mut self.data);
214 let elem_cap = data.capacity();
215 if mem::size_of::<T>() == 0 || elem_cap == 0 {
216 return;
218 }
219 let byte_len = elem_cap * mem::size_of::<T>();
220 if !byte_len.is_multiple_of(BUFFER_ALIGN)
224 || !(data.as_ptr() as usize).is_multiple_of(BUFFER_ALIGN)
225 {
226 return;
227 }
228 data.clear();
231 let n_chunks = byte_len / BUFFER_ALIGN;
232 let ptr = data.as_mut_ptr().cast::<AlignedChunk>();
233 mem::forget(data);
236 let block = unsafe { Vec::<AlignedChunk>::from_raw_parts(ptr, 0, n_chunks) };
241 self.pool.reclaim(n_chunks, block);
242 }
243}
244
245impl<T> Deref for PoolVec<'_, T> {
246 type Target = [T];
247
248 fn deref(&self) -> &[T] {
249 &self.data
250 }
251}
252
253impl<T> DerefMut for PoolVec<'_, T> {
254 fn deref_mut(&mut self) -> &mut [T] {
255 &mut self.data
256 }
257}
258
259impl<T> Extend<T> for PoolVec<'_, T> {
260 fn extend<I: IntoIterator<Item = T>>(&mut self, iter: I) {
261 self.data.extend(iter);
262 }
263}
264
265impl<T: fmt::Debug> fmt::Debug for PoolVec<'_, T> {
266 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
267 f.debug_list().entries(self.data.iter()).finish()
268 }
269}
270
271#[cfg(test)]
272mod tests {
273 use super::*;
274
275 #[test]
276 fn alloc_vec_reserves_capacity_and_starts_empty() {
277 let pool = BufferPool::new();
278 let buffer = pool.alloc_vec::<u64>(16);
279 assert!(buffer.is_empty());
280 assert!(buffer.capacity() >= 16);
281 }
282
283 #[test]
284 fn push_extend_and_deref() {
285 let pool = BufferPool::new();
286 let mut buffer = pool.alloc_vec::<u64>(4);
287 buffer.push(1);
288 buffer.extend_from_slice(&[2, 3]);
289 buffer.extend([4, 5]);
290 assert_eq!(&*buffer, &[1, 2, 3, 4, 5]);
291
292 buffer[0] = 10;
293 assert_eq!(buffer[0], 10);
294
295 buffer.resize(3, 0);
296 assert_eq!(&*buffer, &[10, 2, 3]);
297
298 buffer.clear();
299 assert!(buffer.is_empty());
300 }
301
302 #[test]
303 fn buffers_are_aligned_and_sized_to_a_power_of_two_byte_length() {
304 let pool = BufferPool::new();
305 let buffer = pool.alloc_vec::<u64>(10);
306 assert_eq!(buffer.capacity(), 16);
308 assert!((buffer.as_ptr() as usize).is_multiple_of(BUFFER_ALIGN));
309 assert_eq!(pool.alloc_vec::<u64>(1).capacity(), BUFFER_ALIGN / size_of::<u64>());
311 }
312
313 #[test]
314 fn freed_block_is_recycled_for_a_matching_request() {
315 let pool = BufferPool::new();
316
317 let addr = {
318 let buffer = pool.alloc_vec::<u64>(10);
319 buffer.as_ptr() as usize
320 };
321 let mut buffer = pool.alloc_vec::<u64>(10);
324 assert_eq!(buffer.as_ptr() as usize, addr);
325 assert!(buffer.is_empty());
326 buffer.extend_from_slice(&[1, 2, 3]);
327 assert_eq!(&*buffer, &[1, 2, 3]);
328 }
329
330 #[test]
331 fn a_block_freed_by_one_type_is_reused_by_another_of_the_same_byte_size() {
332 let pool = BufferPool::new();
333 let addr = {
336 let buffer = pool.alloc_vec::<u64>(8);
337 buffer.as_ptr() as usize
338 };
339 let buffer = pool.alloc_vec::<u8>(64);
340 assert_eq!(buffer.as_ptr() as usize, addr);
341 }
342
343 #[test]
344 fn distinct_sizes_do_not_share_blocks() {
345 let pool = BufferPool::new();
346 let small = {
347 let buffer = pool.alloc_vec::<u64>(4);
348 buffer.as_ptr() as usize
349 };
350 let big = pool.alloc_vec::<u64>(64);
352 assert_ne!(big.as_ptr() as usize, small);
353 }
354
355 #[test]
356 fn free_list_holds_multiple_blocks_of_the_same_size() {
357 let pool = BufferPool::new();
358
359 let (addr_a, addr_b) = {
360 let a = pool.alloc_vec::<u64>(8);
361 let b = pool.alloc_vec::<u64>(8);
362 assert_ne!(a.as_ptr() as usize, b.as_ptr() as usize);
363 (a.as_ptr() as usize, b.as_ptr() as usize)
364 };
365
366 let c = pool.alloc_vec::<u64>(8);
368 let d = pool.alloc_vec::<u64>(8);
369 let reused = [c.as_ptr() as usize, d.as_ptr() as usize];
370 assert!(reused.contains(&addr_a));
371 assert!(reused.contains(&addr_b));
372 }
373}