meowalloc 0.1.1

Toy allocator written in pure rust, with concurrency in mind
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)
    }

    /// Attempt to deallocate a block of memory matching 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() };
           
            // The other index now points to what would be the last element before the merge
            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));

    // Memory is writeable and readable
    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() };
    // Layout:
    //   b1      b2           b3                   b4              b0 
    // | 256 |   512   |     1024     |           2048           | 256 |
    // 0    256       768            1792                       3840  4096
    
    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")
    };
    // Layout:
    //              merged                         b4              b0
    // |             1792             |           2048           | 256 |
    // 0                             1792                       3840  4096
   
    assert_eq!(sblock.blocks.0[merged].ptr, init_ptr);
    assert_eq!(sblock.blocks.0[merged].size.get(), 1792);

    // b4 stayed after merge
    assert_eq!(
        sblock.blocks.0.iter()
            .find(|block| block.ptr == unsafe { init_ptr.add(1792) })
            .unwrap().size.get(),
        2048
    );
    
    // b0 stayed after merge
    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);
    // Bitmask state:
    // 0b11110111...
    
    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));
}