#![allow(unsafe_code)]
use std::alloc::{GlobalAlloc, Layout, System};
#[derive(Clone, Copy)]
pub struct Sinks {
alloc: fn(u64),
realloc: fn(u64),
dealloc: fn(),
}
impl Sinks {
#[must_use]
pub const fn noop() -> Self {
Self { alloc: |_| {}, realloc: |_| {}, dealloc: || {} }
}
#[must_use]
pub const unsafe fn new(alloc: fn(u64), realloc: fn(u64), dealloc: fn()) -> Self {
Self { alloc, realloc, dealloc }
}
}
pub struct CountingAlloc<A> {
inner: A,
sinks: Sinks,
}
impl CountingAlloc<System> {
#[must_use]
pub const fn system(sinks: Sinks) -> Self {
Self { inner: System, sinks }
}
}
impl<A> CountingAlloc<A> {
#[must_use]
pub const fn wrapping(inner: A, sinks: Sinks) -> Self {
Self { inner, sinks }
}
}
pub(super) const fn fdu_sinks() -> Sinks {
unsafe { Sinks::new(super::record_alloc, super::record_realloc, super::record_dealloc) }
}
unsafe impl<A: GlobalAlloc> GlobalAlloc for CountingAlloc<A> {
#[inline]
unsafe fn alloc(&self, layout: Layout) -> *mut u8 {
(self.sinks.alloc)(u64::try_from(layout.size()).unwrap_or(u64::MAX));
unsafe { self.inner.alloc(layout) }
}
#[inline]
unsafe fn dealloc(&self, ptr: *mut u8, layout: Layout) {
(self.sinks.dealloc)();
unsafe { self.inner.dealloc(ptr, layout) }
}
#[inline]
unsafe fn realloc(&self, ptr: *mut u8, layout: Layout, new_size: usize) -> *mut u8 {
let growth = u64::try_from(new_size.saturating_sub(layout.size())).unwrap_or(u64::MAX);
(self.sinks.realloc)(growth);
unsafe { self.inner.realloc(ptr, layout, new_size) }
}
#[inline]
unsafe fn alloc_zeroed(&self, layout: Layout) -> *mut u8 {
(self.sinks.alloc)(u64::try_from(layout.size()).unwrap_or(u64::MAX));
unsafe { self.inner.alloc_zeroed(layout) }
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicU64, Ordering};
use super::*;
static ALLOCS: AtomicU64 = AtomicU64::new(0);
static BYTES: AtomicU64 = AtomicU64::new(0);
static FREES: AtomicU64 = AtomicU64::new(0);
fn alloc(size: u64) {
ALLOCS.fetch_add(1, Ordering::Relaxed);
BYTES.fetch_add(size, Ordering::Relaxed);
}
fn realloc(growth: u64) {
BYTES.fetch_add(growth, Ordering::Relaxed);
}
fn dealloc() {
FREES.fetch_add(1, Ordering::Relaxed);
}
fn test_sinks() -> Sinks {
unsafe { Sinks::new(alloc, realloc, dealloc) }
}
#[test]
fn forwards_memory_intact_and_counts_it() {
let _serial = crate::counters::test_serial();
let allocator = CountingAlloc::system(test_sinks());
let layout = Layout::from_size_align(64, 8).expect("valid layout");
ALLOCS.store(0, Ordering::Relaxed);
FREES.store(0, Ordering::Relaxed);
BYTES.store(0, Ordering::Relaxed);
unsafe {
let pointer = allocator.alloc(layout);
assert!(!pointer.is_null());
pointer.write_bytes(0xAB, 64);
assert_eq!(pointer.read(), 0xAB);
allocator.dealloc(pointer, layout);
}
assert_eq!(ALLOCS.load(Ordering::Relaxed), 1);
assert_eq!(FREES.load(Ordering::Relaxed), 1);
assert_eq!(BYTES.load(Ordering::Relaxed), 64);
}
#[test]
fn noop_sinks_are_const_constructible() {
const ALLOCATOR: CountingAlloc<System> = CountingAlloc::system(Sinks::noop());
let layout = Layout::from_size_align(16, 8).expect("valid layout");
unsafe {
let pointer = ALLOCATOR.alloc(layout);
assert!(!pointer.is_null());
ALLOCATOR.dealloc(pointer, layout);
}
}
}