use std::ops::Range;
use std::sync::Arc;
use bytes::Bytes;
use memmap2::MmapMut;
use parking_lot::RwLock;
use super::MappedFileError;
use super::MappedFileResult;
#[derive(Debug, Clone)]
pub struct MappedBuffer {
mmap: Arc<RwLock<MmapMut>>,
offset: usize,
len: usize,
}
impl MappedBuffer {
pub fn new(mmap: Arc<RwLock<MmapMut>>, offset: usize, len: usize) -> MappedFileResult<Self> {
let mmap_guard = mmap.read();
let mmap_len = mmap_guard.len();
drop(mmap_guard);
if offset.checked_add(len).is_none_or(|end| end > mmap_len) {
return Err(MappedFileError::out_of_bounds(offset, len, mmap_len as u64));
}
Ok(Self { mmap, offset, len })
}
pub fn write(&self, offset: usize, data: &[u8]) -> MappedFileResult<()> {
if offset.checked_add(data.len()).is_none_or(|end| end > self.len) {
return Err(MappedFileError::out_of_bounds(
self.offset + offset,
data.len(),
(self.offset + self.len) as u64,
));
}
let mut mmap = self.mmap.write();
let start = self.offset + offset;
let end = start + data.len();
mmap[start..end].copy_from_slice(data);
Ok(())
}
pub fn read(&self, range: Range<usize>) -> MappedFileResult<Bytes> {
if range.end > self.len {
return Err(MappedFileError::out_of_bounds(
self.offset + range.start,
range.len(),
(self.offset + self.len) as u64,
));
}
let mmap = self.mmap.read();
let start = self.offset + range.start;
let end = self.offset + range.end;
Ok(Bytes::copy_from_slice(&mmap[start..end]))
}
pub fn read_zero_copy(&self, range: Range<usize>) -> MappedFileResult<Bytes> {
if range.end > self.len {
return Err(MappedFileError::out_of_bounds(
self.offset + range.start,
range.len(),
(self.offset + self.len) as u64,
));
}
let mmap = self.mmap.read();
let start = self.offset + range.start;
let end = self.offset + range.end;
let size = end - start;
let slice = &mmap[start..end];
let is_aligned = (slice.as_ptr() as usize) % 64 == 0;
if (8192..=65536).contains(&size) && is_aligned {
let mut vec = vec![0u8; size];
unsafe {
std::ptr::copy_nonoverlapping(slice.as_ptr(), vec.as_mut_ptr(), size);
}
Ok(Bytes::from(vec))
} else {
Ok(Bytes::copy_from_slice(slice))
}
}
pub fn batch_write<'a, I>(&self, writes: I) -> MappedFileResult<usize>
where
I: IntoIterator<Item = (usize, &'a [u8])>,
{
let mut mmap = self.mmap.write();
let mut total_written = 0;
for (offset, data) in writes {
if offset.checked_add(data.len()).is_none_or(|end| end > self.len) {
return Err(MappedFileError::out_of_bounds(
self.offset + offset,
data.len(),
(self.offset + self.len) as u64,
));
}
let start = self.offset + offset;
let end = start + data.len();
mmap[start..end].copy_from_slice(data);
total_written += data.len();
}
Ok(total_written)
}
#[inline]
pub fn len(&self) -> usize {
self.len
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len == 0
}
#[inline]
pub fn offset(&self) -> usize {
self.offset
}
pub fn flush(&self) -> MappedFileResult<()> {
let mmap = self.mmap.write();
mmap.flush().map_err(|e| MappedFileError::FlushFailed(e.to_string()))
}
pub fn flush_range(&self, range: Range<usize>) -> MappedFileResult<()> {
if range.end > self.len {
return Err(MappedFileError::out_of_bounds(
self.offset + range.start,
range.len(),
(self.offset + self.len) as u64,
));
}
let mmap = self.mmap.write();
let start = self.offset + range.start;
let end = self.offset + range.end;
mmap.flush_range(start, end - start)
.map_err(|e| MappedFileError::FlushFailed(e.to_string()))
}
pub fn get_mmap(&self) -> Arc<RwLock<MmapMut>> {
Arc::clone(&self.mmap)
}
}
#[cfg(test)]
mod tests {
use std::io::Write as IoWrite;
use tempfile::NamedTempFile;
use super::*;
fn create_test_mmap(size: usize) -> Arc<RwLock<MmapMut>> {
let mut file = NamedTempFile::new().unwrap();
file.write_all(&vec![0u8; size]).unwrap();
file.flush().unwrap();
let file = file.reopen().unwrap();
let mmap = unsafe { MmapMut::map_mut(&file).unwrap() };
Arc::new(RwLock::new(mmap))
}
#[test]
fn test_new_valid_bounds() {
let mmap = create_test_mmap(1024);
let buffer = MappedBuffer::new(mmap, 0, 512);
assert!(buffer.is_ok());
}
#[test]
fn test_new_invalid_bounds() {
let mmap = create_test_mmap(1024);
let buffer = MappedBuffer::new(mmap, 512, 1024);
assert!(buffer.is_err());
}
#[test]
fn test_write_read() {
let mmap = create_test_mmap(1024);
let buffer = MappedBuffer::new(mmap, 0, 1024).unwrap();
buffer.write(0, b"Hello, World!").unwrap();
let data = buffer.read(0..13).unwrap();
assert_eq!(&data[..], b"Hello, World!");
}
#[test]
fn test_write_out_of_bounds() {
let mmap = create_test_mmap(1024);
let buffer = MappedBuffer::new(mmap, 0, 100).unwrap();
let result = buffer.write(90, &[0u8; 20]);
assert!(result.is_err());
}
#[test]
fn test_batch_write() {
let mmap = create_test_mmap(1024);
let buffer = MappedBuffer::new(mmap, 0, 1024).unwrap();
let writes = vec![(0, b"Header" as &[u8]), (6, b"Body"), (10, b"Footer")];
let total = buffer.batch_write(writes).unwrap();
assert_eq!(total, 16);
let data = buffer.read(0..16).unwrap();
assert_eq!(&data[0..6], b"Header");
assert_eq!(&data[6..10], b"Body");
assert_eq!(&data[10..16], b"Footer");
}
#[test]
fn test_zero_copy_read() {
let mmap = create_test_mmap(1024);
let buffer = MappedBuffer::new(mmap, 0, 1024).unwrap();
buffer.write(0, b"Zero Copy Test").unwrap();
let data = buffer.read_zero_copy(0..14).unwrap();
assert_eq!(&data[..], b"Zero Copy Test");
}
}