use std::cell::{Cell, RefCell};
use std::mem::{align_of, size_of};
use std::ptr::NonNull;
use std::slice;
const BLOCK_SIZE: usize = 64 * 1024;
const MAX_FREE_BLOCKS: usize = 16;
const ALIGNMENT: usize = 8;
pub struct MemoryPool {
current_block: RefCell<Option<Block>>,
free_blocks: RefCell<Vec<Block>>,
total_allocated: Cell<usize>,
total_used: Cell<usize>,
}
struct Block {
memory: Vec<u8>,
pos: usize,
}
impl Block {
fn new(size: usize) -> Self {
Block {
memory: vec![0u8; size],
pos: 0,
}
}
#[allow(dead_code)]
fn available(&self) -> usize {
self.memory.len() - self.pos
}
fn allocate(&mut self, size: usize, align: usize) -> Option<NonNull<u8>> {
let aligned_pos = (self.pos + align - 1) & !(align - 1);
let end_pos = aligned_pos + size;
if end_pos <= self.memory.len() {
let ptr = unsafe { self.memory.as_mut_ptr().add(aligned_pos) };
self.pos = end_pos;
NonNull::new(ptr)
} else {
None
}
}
fn reset(&mut self) {
self.pos = 0;
}
}
unsafe impl Send for MemoryPool {}
unsafe impl Sync for MemoryPool {}
impl MemoryPool {
pub fn new() -> Self {
MemoryPool {
current_block: RefCell::new(None),
free_blocks: RefCell::new(Vec::with_capacity(MAX_FREE_BLOCKS)),
total_allocated: Cell::new(0),
total_used: Cell::new(0),
}
}
pub fn allocate(&self, size: usize) -> Option<NonNull<u8>> {
self.allocate_aligned(size, ALIGNMENT)
}
pub fn allocate_aligned(&self, size: usize, align: usize) -> Option<NonNull<u8>> {
if let Some(ref mut block) = *self.current_block.borrow_mut() {
if let Some(ptr) = block.allocate(size, align) {
self.total_used.set(self.total_used.get() + size);
return Some(ptr);
}
}
let block_size = size.max(BLOCK_SIZE);
let mut new_block = self.get_or_create_block(block_size);
let ptr = new_block.allocate(size, align)?;
self.total_used.set(self.total_used.get() + size);
if let Some(old_block) = self.current_block.borrow_mut().take() {
self.store_free_block(old_block);
}
*self.current_block.borrow_mut() = Some(new_block);
Some(ptr)
}
pub fn allocate_str<'a>(&self, s: &str) -> Option<&'a str> {
let bytes = s.as_bytes();
let ptr = self.allocate(bytes.len())?;
unsafe {
std::ptr::copy_nonoverlapping(bytes.as_ptr(), ptr.as_ptr(), bytes.len());
let slice = slice::from_raw_parts(ptr.as_ptr(), bytes.len());
std::str::from_utf8_unchecked(slice).into()
}
}
pub fn allocate_copy<'a, T: Copy>(&self, value: &T) -> Option<&'a T> {
let size = size_of::<T>();
let align = align_of::<T>();
let ptr = self.allocate_aligned(size, align)?;
unsafe {
std::ptr::write(ptr.as_ptr() as *mut T, *value);
Some(&*(ptr.as_ptr() as *const T))
}
}
pub fn allocate_slice<'a, T: Copy>(&self, slice: &[T]) -> Option<&'a [T]> {
if slice.is_empty() {
return Some(unsafe { slice::from_raw_parts(slice.as_ptr(), 0) });
}
let size = size_of_val(slice);
let align = align_of::<T>();
let ptr = self.allocate_aligned(size, align)?;
unsafe {
std::ptr::copy_nonoverlapping(slice.as_ptr(), ptr.as_ptr() as *mut T, slice.len());
Some(slice::from_raw_parts(ptr.as_ptr() as *const T, slice.len()))
}
}
pub fn reset(&self) {
if let Some(ref mut block) = *self.current_block.borrow_mut() {
block.reset();
}
for block in self.free_blocks.borrow_mut().iter_mut() {
block.reset();
}
self.total_used.set(0);
}
pub fn stats(&self) -> MemoryPoolStats {
MemoryPoolStats {
total_allocated: self.total_allocated.get(),
total_used: self.total_used.get(),
num_blocks: self.free_blocks.borrow().len()
+ if self.current_block.borrow().is_some() {
1
} else {
0
},
}
}
fn get_or_create_block(&self, size: usize) -> Block {
let mut free_blocks = self.free_blocks.borrow_mut();
let mut suitable_index = None;
for (i, block) in free_blocks.iter().enumerate() {
if block.memory.len() >= size {
suitable_index = Some(i);
break;
}
}
if let Some(index) = suitable_index {
let mut block = free_blocks.swap_remove(index);
block.reset();
return block;
}
self.total_allocated.set(self.total_allocated.get() + size);
Block::new(size)
}
fn store_free_block(&self, mut block: Block) {
block.reset();
let mut free_blocks = self.free_blocks.borrow_mut();
if free_blocks.len() < MAX_FREE_BLOCKS {
free_blocks.push(block);
}
}
}
impl Default for MemoryPool {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Copy)]
pub struct MemoryPoolStats {
pub total_allocated: usize,
pub total_used: usize,
pub num_blocks: usize,
}
impl MemoryPoolStats {
pub fn utilization(&self) -> f32 {
if self.total_allocated == 0 {
0.0
} else {
(self.total_used as f32 / self.total_allocated as f32) * 100.0
}
}
}
pub struct ScopedMemoryPool<'a> {
pool: MemoryPool,
_phantom: std::marker::PhantomData<&'a ()>,
}
impl<'a> ScopedMemoryPool<'a> {
pub fn new() -> Self {
ScopedMemoryPool {
pool: MemoryPool::new(),
_phantom: std::marker::PhantomData,
}
}
pub fn allocate_str(&self, s: &str) -> Option<&'a str> {
self.pool.allocate_str(s)
}
pub fn allocate_copy<T: Copy>(&self, value: &T) -> Option<&'a T> {
self.pool.allocate_copy(value)
}
pub fn allocate_slice<T: Copy>(&self, slice: &[T]) -> Option<&'a [T]> {
self.pool.allocate_slice(slice)
}
pub fn reset(&self) {
self.pool.reset()
}
pub fn stats(&self) -> MemoryPoolStats {
self.pool.stats()
}
}
impl<'a> Default for ScopedMemoryPool<'a> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_allocation() {
let pool = MemoryPool::new();
let ptr1 = pool.allocate(100).unwrap();
let ptr2 = pool.allocate(200).unwrap();
assert_ne!(ptr1.as_ptr(), ptr2.as_ptr());
let stats = pool.stats();
assert_eq!(stats.total_used, 300);
}
#[test]
fn test_string_allocation() {
let pool = MemoryPool::new();
let s1 = "Hello, world!";
let s2 = "Another string";
let allocated1 = pool.allocate_str(s1).unwrap();
let allocated2 = pool.allocate_str(s2).unwrap();
assert_eq!(allocated1, s1);
assert_eq!(allocated2, s2);
assert_ne!(allocated1.as_ptr(), allocated2.as_ptr());
}
#[test]
fn test_reset() {
let pool = MemoryPool::new();
pool.allocate(1000).unwrap();
let stats_before = pool.stats();
assert_eq!(stats_before.total_used, 1000);
pool.reset();
let stats_after = pool.stats();
assert_eq!(stats_after.total_used, 0);
assert_eq!(stats_after.total_allocated, stats_before.total_allocated);
}
#[test]
fn test_large_allocation() {
let pool = MemoryPool::new();
let large_size = BLOCK_SIZE * 2;
let _ptr = pool.allocate(large_size).unwrap();
let stats = pool.stats();
assert_eq!(stats.total_used, large_size);
}
#[test]
fn test_scoped_pool() {
let pool = ScopedMemoryPool::new();
let s = "Test string";
let allocated = pool.allocate_str(s).unwrap();
assert_eq!(allocated, s);
}
}