#![deny(missing_docs)]
#![deny(clippy::indexing_slicing)]
#![deny(clippy::unwrap_used)]
#![deny(clippy::expect_used)]
#![deny(clippy::panic)]
#![cfg_attr(
test,
allow(
clippy::expect_used,
clippy::indexing_slicing,
clippy::panic,
clippy::unwrap_used
)
)]
use std::alloc::{GlobalAlloc, Layout, System};
use std::cell::Cell;
pub const LARGEST: usize = 4_096;
const GRAIN: usize = 16;
const CLASSES: usize = LARGEST / GRAIN + 1;
const PER_CLASS_BYTES: usize = 64 << 10;
const PER_CLASS_BLOCKS: usize = 1_024;
const CAPS: [usize; CLASSES] = caps();
#[allow(clippy::indexing_slicing)]
const fn caps() -> [usize; CLASSES] {
let mut caps = [1usize; CLASSES];
let mut class = 0usize;
while class < CLASSES {
let size = if class * GRAIN > GRAIN {
class * GRAIN
} else {
GRAIN
};
let mut cap = PER_CLASS_BYTES / size;
if cap > PER_CLASS_BLOCKS {
cap = PER_CLASS_BLOCKS;
}
if cap < 1 {
cap = 1;
}
caps[class] = cap;
class += 1;
}
caps
}
#[inline]
fn per_class(class: usize) -> usize {
CAPS.get(class).copied().unwrap_or(1)
}
thread_local! {
static HEADS: [Cell<*mut u8>; CLASSES] =
const { [const { Cell::new(std::ptr::null_mut()) }; CLASSES] };
static HELD: [Cell<usize>; CLASSES] = const { [const { Cell::new(0) }; CLASSES] };
}
#[inline]
fn class_of(layout: Layout) -> Option<usize> {
if layout.size() > LARGEST
|| layout.align() > GRAIN
|| layout.size() < core::mem::size_of::<*mut u8>()
{
return None;
}
Some(layout.size().div_ceil(GRAIN))
}
#[inline]
fn layout_of(class: usize) -> Layout {
Layout::from_size_align(class.saturating_mul(GRAIN).max(GRAIN), GRAIN)
.unwrap_or_else(|_| Layout::new::<u128>())
}
pub struct Pooled;
unsafe impl GlobalAlloc for Pooled {
#[inline]
unsafe fn alloc(&self, layout: Layout) -> *mut u8 {
let Some(class) = class_of(layout) else {
return unsafe { System.alloc(layout) };
};
let taken = HEADS
.try_with(|heads| {
let Some(head) = heads.get(class) else {
return std::ptr::null_mut();
};
let block = head.get();
if block.is_null() {
return std::ptr::null_mut();
}
let next = unsafe { block.cast::<*mut u8>().read() };
head.set(next);
let _ = HELD.try_with(|held| {
if let Some(count) = held.get(class) {
count.set(count.get().saturating_sub(1));
}
});
block
})
.unwrap_or(std::ptr::null_mut());
if !taken.is_null() {
return taken;
}
unsafe { System.alloc(layout_of(class)) }
}
#[inline]
unsafe fn dealloc(&self, pointer: *mut u8, layout: Layout) {
let Some(class) = class_of(layout) else {
return unsafe { System.dealloc(pointer, layout) };
};
let kept = HELD
.try_with(|held| {
let Some(count) = held.get(class) else {
return false;
};
if count.get() >= per_class(class) {
return false;
}
HEADS
.try_with(|heads| {
let Some(head) = heads.get(class) else {
return false;
};
unsafe { pointer.cast::<*mut u8>().write(head.get()) };
head.set(pointer);
count.set(count.get().saturating_add(1));
true
})
.unwrap_or(false)
})
.unwrap_or(false);
if kept {
return;
}
unsafe { System.dealloc(pointer, layout_of(class)) };
}
#[inline]
unsafe fn realloc(&self, pointer: *mut u8, layout: Layout, new_size: usize) -> *mut u8 {
let old = class_of(layout);
let new = Layout::from_size_align(new_size, layout.align())
.ok()
.and_then(class_of);
if let (Some(old), Some(new)) = (old, new) {
if old == new {
return pointer;
}
}
unsafe {
let Ok(wanted) = Layout::from_size_align(new_size, layout.align()) else {
return std::ptr::null_mut();
};
let fresh = self.alloc(wanted);
if !fresh.is_null() {
std::ptr::copy_nonoverlapping(pointer, fresh, layout.size().min(new_size));
self.dealloc(pointer, layout);
}
fresh
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_classes_are_the_ones_documented() {
assert_eq!(class_of(Layout::from_size_align(8, 8).unwrap()), Some(1));
assert_eq!(class_of(Layout::from_size_align(16, 16).unwrap()), Some(1));
assert_eq!(class_of(Layout::from_size_align(17, 8).unwrap()), Some(2));
assert_eq!(
class_of(Layout::from_size_align(LARGEST, 16).unwrap()),
Some(CLASSES - 1)
);
assert_eq!(
class_of(Layout::from_size_align(LARGEST + 1, 8).unwrap()),
None
);
assert_eq!(class_of(Layout::from_size_align(32, 32).unwrap()), None);
assert_eq!(class_of(Layout::from_size_align(4, 4).unwrap()), None);
}
#[test]
fn a_class_block_covers_every_request_in_it() {
for size in 8..=LARGEST {
let Some(layout) = Layout::from_size_align(size, 8).ok() else {
continue;
};
let Some(class) = class_of(layout) else {
continue;
};
let block = layout_of(class);
assert!(block.size() >= layout.size(), "size {size}");
assert!(block.align() >= layout.align(), "size {size}");
}
}
#[test]
fn a_freed_block_is_the_one_handed_back() {
let layout = Layout::from_size_align(64, 8).expect("a layout");
unsafe {
let first = Pooled.alloc(layout);
assert!(!first.is_null());
Pooled.dealloc(first, layout);
let second = Pooled.alloc(layout);
assert_eq!(first, second, "the freed block was not recycled");
Pooled.dealloc(second, layout);
}
}
#[test]
fn recycling_stays_inside_the_block() {
let layout = Layout::from_size_align(16, 8).expect("a layout");
unsafe {
let guard = Pooled.alloc(layout);
let block = Pooled.alloc(layout);
std::ptr::write_bytes(guard, 0xAB, layout.size());
Pooled.dealloc(block, layout);
let again = Pooled.alloc(layout);
assert_eq!(block, again);
for at in 0..layout.size() {
assert_eq!(guard.add(at).read(), 0xAB, "the link ran past its block");
}
Pooled.dealloc(again, layout);
Pooled.dealloc(guard, layout);
}
}
#[test]
fn realloc_keeps_the_bytes() {
unsafe {
let small = Layout::from_size_align(16, 8).expect("a layout");
let block = Pooled.alloc(small);
std::ptr::write_bytes(block, 0x5A, small.size());
let inside = Pooled.realloc(block, small, 12);
assert_eq!(inside, block, "a realloc inside one class moved the block");
let across = Pooled.realloc(inside, small, 200);
assert!(!across.is_null());
for at in 0..small.size() {
assert_eq!(across.add(at).read(), 0x5A, "realloc lost a byte");
}
Pooled.dealloc(across, Layout::from_size_align(200, 8).expect("a layout"));
}
}
}