use bytemuck::{Pod, Zeroable};
use std::{
alloc::{alloc, dealloc, Layout},
ptr,
};
use crate::error::MesoError;
pub type LogState = (*mut u8, u64);
#[derive(Debug)]
pub struct Mnemosyne {
arena: *mut u8,
size: usize,
write: usize,
state: LogState,
tape: Vec<LogState>,
current_writes: Vec<LogState>,
allocations: Vec<(*mut u8, Layout)>,
}
impl Mnemosyne {
pub fn initialize(size: usize) -> Self {
let layout = Layout::from_size_align(size, 8).unwrap();
let arena = unsafe { alloc(layout) };
let state = (ptr::null_mut::<u8>(), 0u64);
let tape = Vec::new();
let current_writes = Vec::new();
Self {
arena,
size,
write: 0,
state,
tape,
current_writes,
allocations: vec![(arena, layout)],
}
}
pub fn write<T: Pod + Zeroable + 'static>(&mut self, state: T, time: u64) {
if self.arena.is_null() {
let layout = Layout::from_size_align(self.size, 8).unwrap();
unsafe {
let arena = alloc(layout);
self.arena = arena;
self.allocations.push((arena, layout));
}
}
let bytes: &[u8] = bytemuck::bytes_of(&state);
let size = bytes.len();
let align = std::mem::align_of_val(&state);
let offset = (self.write + align - 1) & !(align - 1);
let mut end = offset + size;
if end > self.size {
self.flush(true);
let offset = (align - 1) & !(align - 1);
end = offset + size;
if end > self.size {
unsafe {
let layout = Layout::from_size_align(size, align).unwrap();
let ptr = alloc(layout);
self.allocations.push((ptr, layout));
let dst = std::slice::from_raw_parts_mut(ptr, size);
let src = std::slice::from_raw_parts(&state as *const T as *const u8, size);
dst.copy_from_slice(src);
let _ = state;
self.state = (dst.as_mut_ptr(), time);
self.current_writes.push(self.state);
}
return;
}
}
unsafe {
let dst = std::slice::from_raw_parts_mut(self.arena.add(offset), size);
let src = std::slice::from_raw_parts(&state as *const T as *const u8, size);
dst.copy_from_slice(src);
let _ = state;
self.state = (dst.as_mut_ptr(), time);
self.write = end;
self.current_writes.push(self.state);
}
}
fn flush(&mut self, reset: bool) {
if self.write != 0 {
let writes = std::mem::take(&mut self.current_writes);
self.tape.extend(writes);
if reset {
let layout = Layout::from_size_align(self.size, 8).unwrap();
let arena = unsafe { alloc(layout) };
self.arena = arena;
self.allocations.push((arena, layout));
} else {
self.arena = ptr::null_mut();
}
self.write = 0;
}
}
pub fn read_state<T: Pod + Zeroable + 'static>(&self) -> Result<&T, MesoError> {
let (ptr, _) = self.state;
if ptr.is_null() {
return Err(MesoError::UninitializedState);
}
let out = unsafe { &*(ptr as *const T) };
Ok(out)
}
pub fn read_state_mut<T: Pod + Zeroable + 'static>(&mut self) -> Result<&mut T, MesoError> {
let (ptr, _) = self.state;
if ptr.is_null() {
return Err(MesoError::UninitializedState);
}
let out = unsafe { &mut *(ptr as *mut T) };
Ok(out)
}
pub fn read_tape<T: Pod + Zeroable + 'static>(&self) -> Vec<(&T, u64)> {
let mut out = Vec::new();
for (ptr, time) in &self.tape {
unsafe {
let data = &*(*ptr as *const T);
out.push((data, *time))
}
}
out
}
pub fn read_tape_mut<T: Pod + Zeroable + 'static>(&mut self) -> Vec<(&mut T, u64)> {
let mut out = Vec::new();
for (ptr, time) in &self.tape {
unsafe {
let data = &mut *(*ptr as *mut T);
out.push((data, *time))
}
}
out
}
pub fn cleanup<T: Pod + Zeroable + 'static>(&mut self) -> Vec<(T, u64)> {
let mut out = Vec::new();
self.flush(false);
for (ptr, time) in &self.tape {
unsafe {
let data = ptr::read(*ptr as *mut T);
out.push((data, *time));
}
}
for (i, layout) in &self.allocations {
unsafe { dealloc(*i, *layout) };
}
self.write = 0;
self.state = (ptr::null_mut(), 0);
self.tape.clear();
self.current_writes.clear();
self.allocations.clear();
out
}
}
impl Drop for Mnemosyne {
fn drop(&mut self) {
self.flush(false);
for (i, layout) in &self.allocations {
unsafe { dealloc(*i, *layout) };
}
self.write = 0;
self.state = (ptr::null_mut(), 0);
self.tape.clear();
self.current_writes.clear();
self.allocations.clear();
}
}
#[cfg(test)]
mod tests {
use super::*; use bytemuck::{Pod, Zeroable};
#[derive(Copy, Clone, Debug, PartialEq)]
#[repr(C)] struct MyState {
x: u32,
y: f32,
z: u64,
}
unsafe impl Pod for MyState {}
unsafe impl Zeroable for MyState {}
#[test]
fn test_initialize() {
let size = 1024;
let mnemosyne = Mnemosyne::initialize(size);
assert_eq!(mnemosyne.size, size);
assert_eq!(mnemosyne.write, 0);
assert!(!mnemosyne.arena.is_null());
assert_eq!(mnemosyne.state, (ptr::null_mut(), 0));
assert!(mnemosyne.tape.is_empty());
assert!(mnemosyne.current_writes.is_empty());
assert_eq!(mnemosyne.allocations.len(), 1); }
#[test]
fn test_write_primitive_and_read_state() {
let size = 64; let mut mnemosyne = Mnemosyne::initialize(size);
let val_u32 = 12345u32;
let time_u32 = 100u64;
mnemosyne.write(val_u32, time_u32);
assert_eq!(mnemosyne.read_state::<u32>().unwrap(), &val_u32);
assert_eq!(mnemosyne.state.1, time_u32);
assert_eq!(mnemosyne.current_writes.len(), 1);
assert_eq!(mnemosyne.write, std::mem::size_of::<u32>());
let val_f32 = 3.10f32;
let time_f32 = 200u64;
mnemosyne.write(val_f32, time_f32);
assert_eq!(mnemosyne.read_state::<f32>().unwrap(), &val_f32);
assert_eq!(mnemosyne.state.1, time_f32);
assert_eq!(mnemosyne.current_writes.len(), 2);
}
#[test]
fn test_write_struct_and_read_state() {
let size = 64;
let mut mnemosyne = Mnemosyne::initialize(size);
let state = MyState {
x: 10,
y: 20.5,
z: 300,
};
let time = 500u64;
mnemosyne.write(state, time);
assert_eq!(mnemosyne.read_state::<MyState>().unwrap(), &state);
assert_eq!(mnemosyne.state.1, time);
assert_eq!(mnemosyne.current_writes.len(), 1);
assert_eq!(mnemosyne.write, std::mem::size_of::<MyState>());
}
#[test]
fn test_read_state_uninitialized() {
let size = 64;
let mnemosyne = Mnemosyne::initialize(size);
assert_eq!(
mnemosyne.read_state::<u32>(),
Err(MesoError::UninitializedState)
);
}
#[test]
fn test_arena_overflow_and_flush_true() {
let size = std::mem::size_of::<u32>() * 2 + 1; let mut mnemosyne = Mnemosyne::initialize(size);
let val1 = 1u32;
let time1 = 100u64;
mnemosyne.write(val1, time1);
let val2 = 2u32;
let time2 = 200u64;
mnemosyne.write(val2, time2);
let val3 = 3u32;
let time3 = 300u64;
mnemosyne.write(val3, time3);
assert_eq!(mnemosyne.tape.len(), 2); assert_eq!(mnemosyne.read_tape::<u32>()[0], (&val1, time1));
assert_eq!(mnemosyne.read_tape::<u32>()[1], (&val2, time2));
assert_eq!(mnemosyne.read_state::<u32>().unwrap(), &val3);
assert_eq!(mnemosyne.state.1, time3);
assert_eq!(mnemosyne.allocations.len(), 2);
assert_eq!(mnemosyne.write, std::mem::align_of::<u32>());
}
#[test]
fn test_write_too_large_for_arena() {
let size = std::mem::size_of::<MyState>() / 2;
let mut mnemosyne = Mnemosyne::initialize(size);
let state = MyState {
x: 111,
y: 22.2,
z: 333,
};
let time = 700u64;
mnemosyne.write(state, time);
assert_eq!(mnemosyne.read_state::<MyState>().unwrap(), &state);
assert_eq!(mnemosyne.state.1, time);
assert_eq!(mnemosyne.current_writes.len(), 1);
assert_eq!(mnemosyne.tape.len(), 0); assert_eq!(mnemosyne.write, 0); assert_eq!(mnemosyne.allocations.len(), 2); }
#[test]
fn test_read_tape() {
let size = 256;
let mut mnemosyne = Mnemosyne::initialize(size);
let s1 = MyState {
x: 1,
y: 1.0,
z: 10,
};
let t1 = 100;
mnemosyne.write(s1, t1);
let s2 = MyState {
x: 2,
y: 2.0,
z: 20,
};
let t2 = 200;
mnemosyne.write(s2, t2);
mnemosyne.flush(false);
let s3 = MyState {
x: 3,
y: 3.0,
z: 30,
};
let t3 = 300;
mnemosyne.write(s3, t3);
let tape_data = mnemosyne.read_tape::<MyState>();
assert_eq!(tape_data.len(), 2);
assert_eq!(tape_data[0], (&s1, t1));
assert_eq!(tape_data[1], (&s2, t2));
assert_eq!(mnemosyne.read_state::<MyState>().unwrap(), &s3);
}
#[test]
fn test_read_tape_mut() {
let size = 256;
let mut mnemosyne = Mnemosyne::initialize(size);
let s1 = MyState {
x: 1,
y: 1.0,
z: 10,
};
let t1 = 100;
mnemosyne.write(s1, t1);
let s2 = MyState {
x: 2,
y: 2.0,
z: 20,
};
let t2 = 200;
mnemosyne.write(s2, t2);
mnemosyne.flush(false);
let mut tape_data_mut = mnemosyne.read_tape_mut::<MyState>();
assert_eq!(tape_data_mut.len(), 2);
tape_data_mut[0].0.x = 111;
tape_data_mut[1].0.y = 222.0;
let tape_data = mnemosyne.read_tape::<MyState>();
assert_eq!(tape_data[0].0.x, 111);
assert_eq!(tape_data[1].0.y, 222.0);
}
#[test]
fn test_cleanup() {
let size = 64;
let mut mnemosyne = Mnemosyne::initialize(size);
let s1 = MyState {
x: 1,
y: 1.0,
z: 10,
};
let t1 = 100;
mnemosyne.write(s1, t1);
let s2 = MyState {
x: 2,
y: 2.0,
z: 20,
};
let t2 = 200;
mnemosyne.write(s2, t2);
let s3 = MyState {
x: 3,
y: 3.0,
z: 30,
}; let t3 = 300;
mnemosyne.write(s3, t3);
let collected_data = mnemosyne.cleanup::<MyState>();
assert_eq!(collected_data.len(), 3);
assert_eq!(collected_data[0], (s1, t1));
assert_eq!(collected_data[1], (s2, t2));
assert_eq!(collected_data[2], (s3, t3));
assert_eq!(mnemosyne.write, 0);
assert_eq!(mnemosyne.state, (ptr::null_mut(), 0));
assert!(mnemosyne.tape.is_empty());
assert!(mnemosyne.current_writes.is_empty());
assert!(mnemosyne.allocations.is_empty()); assert!(mnemosyne.arena.is_null()); }
}