use crate::{send_sync_ptr::SendSyncPtr, traits::Reset};
use eyre::Result;
use memfd::{Memfd, MemfdOptions};
use rustix::mm::{self, MapFlags, ProtFlags};
use std::{
ptr::{self, NonNull},
rc::Rc,
};
#[repr(C)]
pub struct ForkableMemory {
memory: SendSyncPtr<[u8]>,
fd: Rc<Memfd>,
}
impl ForkableMemory {
pub fn new(size: usize) -> Result<Self> {
let mfd = MemfdOptions::default().create(format!("sized-{size}"))?;
mfd.as_file().set_len(size as u64)?;
let ptr = unsafe {
mm::mmap(
ptr::null_mut(),
size,
ProtFlags::READ | ProtFlags::WRITE,
MapFlags::SHARED | MapFlags::NORESERVE,
mfd.as_file(),
0,
)?
};
let memory = std::ptr::slice_from_raw_parts_mut(ptr.cast(), size);
Ok(Self {
memory: NonNull::new(memory).unwrap().into(),
fd: Rc::new(mfd),
})
}
#[inline]
pub fn as_ptr(&self) -> *mut u8 {
self.memory.cast().as_ptr()
}
#[inline]
pub fn as_send_sync_ptr(&self) -> SendSyncPtr<u8> {
self.memory.cast()
}
pub fn fork(&self) -> Result<Self> {
let ptr = unsafe {
mm::mmap(
ptr::null_mut(),
self.memory.len(),
ProtFlags::READ | ProtFlags::WRITE,
MapFlags::PRIVATE | MapFlags::NORESERVE,
self.fd.as_file(),
0,
)?
};
let memory = std::ptr::slice_from_raw_parts_mut(ptr.cast(), self.memory.len());
Ok(Self {
memory: NonNull::new(memory).unwrap().into(),
fd: self.fd.clone(),
})
}
}
impl Reset for ForkableMemory {
fn reset(&mut self) {
let size = self.memory.len();
self.fd.as_file().set_len(0).unwrap();
self.fd.as_file().set_len(size as u64).unwrap();
let ptr = unsafe {
mm::mmap(
self.as_ptr().cast(),
size,
ProtFlags::READ | ProtFlags::WRITE,
MapFlags::SHARED | MapFlags::FIXED | MapFlags::NORESERVE,
self.fd.as_file(),
0,
)
.unwrap()
};
let memory = std::ptr::slice_from_raw_parts_mut(ptr.cast(), size);
self.memory = NonNull::new(memory).unwrap().into();
}
}
impl Drop for ForkableMemory {
fn drop(&mut self) {
unsafe {
let ptr = self.memory.as_ptr().cast();
let len = self.memory.len();
if len == 0 {
return;
}
rustix::mm::munmap(ptr, len).expect("munmap failed");
}
}
}