use kitset::KitSet;
use meowvec::MeowVec;
use core::{alloc::{AllocError, Layout}, num::NonZero, ptr::NonNull};
use crate::{atomic_bitset::WORD_BITS, constants::linux::PAGE_SIZE};
use crate::sys::{mmap_anon, munmap};
use crate::block::ExtractIdx;
use crate::block_store::BlockStore;
pub struct Superblock<const MAX_BLOCKS: usize, const FREE_BITSET_WORDS: usize> {
ptr: NonNull<u8>,
size: NonZero<usize>,
blocks: BlockStore<MAX_BLOCKS>,
free: KitSet<FREE_BITSET_WORDS>,
}
enum MergeStatus {
DidMerge(usize),
DidNotMerge
}
impl<
const MAX_BLOCKS: usize,
const FREE_BITSET_WORDS: usize
> Superblock<MAX_BLOCKS, FREE_BITSET_WORDS>
{
pub fn new(size: NonZero<usize>) -> Result<Self, AllocError> {
const { assert!(FREE_BITSET_WORDS >= MAX_BLOCKS / WORD_BITS) }
let padding = size.get() % PAGE_SIZE;
let size = size.checked_add(padding).ok_or(AllocError)?;
let ptr = mmap_anon(size)?;
let blocks = BlockStore::new(ptr, size);
let mut free = KitSet::<FREE_BITSET_WORDS>::new();
free.set_one(0);
Ok(Self { ptr, size, blocks, free })
}
pub fn alloc(&mut self, layout: Layout) -> Result<NonNull<[u8]>, AllocError> {
if layout.size() == 0 {
return Ok(layout.dangling_ptr().cast_slice(0))
};
let idx = unsafe { self.blocks.create(layout, self.free.iter_ones())? };
self.free.flip(idx);
let ptr = unsafe {
self.blocks.0.get_unchecked(idx).ptr.cast_slice(layout.size())
};
Ok(ptr)
}
pub unsafe fn dealloc(&mut self, ptr: NonNull<u8>) -> Result<(), AllocError> {
let idx = unsafe {
self.blocks
.find_ptr(ptr, self.free.iter_zeros())
.ok_or(AllocError)?
.extract()
};
self.free.flip(idx);
unsafe { self.merge_adjacent_all(idx)? };
Ok(())
}
unsafe fn merge_adjacent_all(&mut self, mut idx: usize) -> Result<usize, AllocError> {
loop {
idx = match unsafe { self.merge_adjacent(idx)? } {
MergeStatus::DidMerge(merged) => merged,
MergeStatus::DidNotMerge => return Ok(idx)
}
}
}
unsafe fn merge_adjacent(&mut self, mut idx: usize) -> Result<MergeStatus, AllocError> {
let free_adjacent: MeowVec<usize, 2> = unsafe {
let block = self.blocks.0.get_unchecked(idx);
self.blocks
.find_adjacent(block, self.free.iter_ones())
.map(|guard| guard.extract())
.collect()
};
let did_merge = ! free_adjacent.is_empty();
for other_idx in free_adjacent {
let old_idx = idx;
let old_last_state = self.free.is_one(self.blocks_len() - 1);
idx = unsafe { self.blocks.merge(idx, other_idx)?.extract() };
self.free.set(
if idx == old_idx { other_idx } else { old_idx },
old_last_state
);
self.free.set_zero(self.blocks_len());
self.free.set_one(idx);
};
if did_merge {
Ok(MergeStatus::DidMerge(idx))
} else {
Ok(MergeStatus::DidNotMerge)
}
}
pub fn blocks_len(&self) -> usize { self.blocks.len() }
pub fn blocks(&self) -> &BlockStore<MAX_BLOCKS> { &self.blocks }
pub fn ptr(&self) -> NonNull<u8> { self.ptr }
pub fn size(&self) -> NonZero<usize> { self.size }
}
impl<
const MAX_BLOCKS: usize,
const FREE_BITSET_WORDS: usize
> Drop for Superblock<MAX_BLOCKS, FREE_BITSET_WORDS> {
fn drop(&mut self) {
unsafe { munmap(self.ptr, self.size) }
}
}
#[cfg(test)]
extern crate std;
#[cfg(test)]
use std::dbg;
#[cfg(test)]
impl<const BLOCKS: usize, const WORDS: usize> Superblock<BLOCKS, WORDS> {
fn divide(&mut self, parts: usize) {
assert!(self.size.get() % parts == 0);
let size = self.size.get() / parts;
for idx in 0..parts {
self.blocks.split(idx, size).unwrap();
self.free.set_one(idx);
}
}
}
#[cfg(test)]
fn superblock() -> Superblock<16, 1> {
Superblock::new(PAGE_SIZE).unwrap()
}
#[test]
fn new() {
let mut sblock = superblock();
let block = &mut sblock.blocks.0[0];
assert_eq!(block.size, PAGE_SIZE);
assert!(sblock.free.is_one(0));
unsafe {
block.ptr.write_bytes(1, PAGE_SIZE.get());
assert_eq!(
block.ptr.cast_slice(PAGE_SIZE.get()).as_ref(),
&[1; PAGE_SIZE.get()]
)
}
}
#[test]
fn alloc() {
let mut sblock = superblock();
let layout = Layout::new::<u128>();
let size = NonZero::new(size_of::<u128>()).unwrap();
let mut ptr = sblock.alloc(layout).unwrap().cast::<u128>();
let num = unsafe { ptr.as_mut() };
*num = u128::MAX;
assert_eq!(*num, u128::MAX);
assert!(
sblock.blocks.0.iter().any(|block| {
(block.ptr.addr() == ptr.addr()) && (block.size >= size)
})
);
}
#[test]
fn merge_adjacent() {
let mut sblock = superblock();
let init_ptr = sblock.blocks.0[0].ptr;
let (b1, free) = unsafe { sblock.blocks.split(0, 256).unwrap().extract() };
let (b2, free) = unsafe { sblock.blocks.split(free, 512).unwrap().extract() };
let (b3, free) = unsafe { sblock.blocks.split(free, 1024).unwrap().extract() };
let (b4, leftover) = unsafe { sblock.blocks.split(free, 2048).unwrap().extract() };
sblock.free.set_one(b1);
sblock.free.set_one(b2);
sblock.free.set_one(b3);
sblock.free.set_one(b4);
sblock.free.set_one(leftover);
let status = unsafe { sblock.merge_adjacent(b2).unwrap() };
let merged = match status {
MergeStatus::DidMerge(idx) => idx,
MergeStatus::DidNotMerge => panic!("should merge")
};
assert_eq!(sblock.blocks.0[merged].ptr, init_ptr);
assert_eq!(sblock.blocks.0[merged].size.get(), 1792);
assert_eq!(
sblock.blocks.0.iter()
.find(|block| block.ptr == unsafe { init_ptr.add(1792) })
.unwrap().size.get(),
2048
);
assert_eq!(
sblock.blocks.0.iter()
.find(|block| block.ptr == unsafe { init_ptr.add(3840) })
.unwrap().size.get(),
256
);
assert_eq!(sblock.free.first_one(), Some(0));
assert_eq!(sblock.free.last_one(), Some(2));
assert_eq!(sblock.free.first_zero(), Some(3));
}
#[test]
fn merge_adjacent_all() {
let mut sblock = superblock();
let init_ptr = sblock.blocks.0[0].ptr;
sblock.divide(8);
sblock.free.set_zero(4);
let merged = unsafe { sblock.merge_adjacent_all(2).unwrap() };
let block = &sblock.blocks.0[merged];
assert_eq!(block.ptr, init_ptr);
assert_eq!(block.size.get(), 2048);
assert!(sblock.free.is_one(merged));
}
#[test]
fn dealloc() {
let mut sblock = superblock();
sblock.blocks.split(0, 2048).unwrap();
let init_ptr = sblock.blocks.0[0].ptr;
let ptr = sblock.blocks.0[1].ptr;
unsafe { sblock.dealloc(ptr).unwrap() };
assert_eq!(sblock.blocks.len(), 1);
assert_eq!(sblock.blocks.0[0].ptr, init_ptr);
assert_eq!(sblock.blocks.0[0].size, PAGE_SIZE);
assert_eq!(sblock.free.first_one(), Some(0));
assert_eq!(sblock.free.last_one(), Some(0));
}