use std::path::{Path, PathBuf};
use async_trait::async_trait;
use memmap2::MmapMut;
use tracing::{debug, warn};
use crate::error::{Aria2Error, Result};
use super::disk_writer::SeekableDiskWriter;
use super::positioned_disk_writer::PositionedDiskWriter;
enum Inner {
Mmap {
#[allow(dead_code)]
file: std::fs::File,
mmap: MmapMut,
},
Fallback(PositionedDiskWriter),
}
pub struct MmapDiskWriter {
inner: Option<Inner>,
path: PathBuf,
total_size: Option<u64>,
opened: bool,
}
impl MmapDiskWriter {
pub fn new(path: &Path, total_size: Option<u64>) -> Self {
Self {
inner: None,
path: path.to_path_buf(),
total_size,
opened: false,
}
}
fn ensure_parent_dirs(&self) -> Result<()> {
if let Some(parent) = self.path.parent()
&& !parent.as_os_str().is_empty()
&& !parent.exists()
{
std::fs::create_dir_all(parent)?;
debug!("Created parent directories for {:?}", self.path);
}
Ok(())
}
fn open_sync(&mut self) -> Result<()> {
self.ensure_parent_dirs()?;
let file = std::fs::OpenOptions::new()
.create(true)
.write(true)
.read(true)
.truncate(false) .open(&self.path)?;
if let Some(size) = self.total_size {
let current_size = file.metadata()?.len();
if current_size == 0 && size > 0 {
file.set_len(size)?;
debug!("Pre-allocated file to {} bytes: {:?}", size, self.path);
}
}
let file_size = file.metadata()?.len();
if file_size == 0 {
warn!(
"File size is 0, cannot mmap: {:?}, using positioned I/O fallback",
self.path
);
drop(file);
let writer = PositionedDiskWriter::new(&self.path, self.total_size);
self.inner = Some(Inner::Fallback(writer));
return Ok(());
}
match unsafe { MmapMut::map_mut(&file) } {
Ok(mmap) => {
debug!(
"Created mmap for {:?}, size: {} bytes",
self.path, file_size
);
self.inner = Some(Inner::Mmap { file, mmap });
}
Err(e) => {
warn!(
"mmap failed for {:?}: {}, using positioned I/O fallback",
self.path, e
);
drop(file);
let writer = PositionedDiskWriter::new(&self.path, self.total_size);
self.inner = Some(Inner::Fallback(writer));
}
}
Ok(())
}
}
#[async_trait]
impl SeekableDiskWriter for MmapDiskWriter {
async fn open(&mut self) -> Result<()> {
if self.opened {
return Ok(());
}
if self.inner.is_none() {
self.open_sync()?;
}
if let Some(Inner::Fallback(ref mut writer)) = self.inner {
writer.open().await?;
}
self.opened = true;
Ok(())
}
async fn write_at(&mut self, offset: u64, data: &[u8]) -> Result<()> {
self.open().await?;
match self.inner.as_mut() {
Some(Inner::Mmap { mmap, .. }) => {
let start = offset as usize;
let end = start
.checked_add(data.len())
.ok_or_else(|| Aria2Error::Io("write offset + length overflow".into()))?;
if end > mmap.len() {
return Err(Aria2Error::Io(format!(
"write at offset {} len {} exceeds mmap size {}",
offset,
data.len(),
mmap.len()
)));
}
mmap[start..end].copy_from_slice(data);
Ok(())
}
Some(Inner::Fallback(writer)) => writer.write_at(offset, data).await,
None => Err(Aria2Error::Io("writer not open".into())),
}
}
async fn write_bytes_at(&mut self, offset: u64, data: bytes::Bytes) -> Result<()> {
self.open().await?;
match self.inner.as_mut() {
Some(Inner::Mmap { mmap, .. }) => {
let start = offset as usize;
let end = start
.checked_add(data.len())
.ok_or_else(|| Aria2Error::Io("write offset + length overflow".into()))?;
if end > mmap.len() {
return Err(Aria2Error::Io(format!(
"write at offset {} len {} exceeds mmap size {}",
offset,
data.len(),
mmap.len()
)));
}
mmap[start..end].copy_from_slice(&data);
Ok(())
}
Some(Inner::Fallback(writer)) => writer.write_bytes_at(offset, data).await,
None => Err(Aria2Error::Io("writer not open".into())),
}
}
async fn read_at(&mut self, offset: u64, buf: &mut [u8]) -> Result<usize> {
self.open().await?;
match self.inner.as_mut() {
Some(Inner::Mmap { mmap, .. }) => {
let start = offset as usize;
if start >= mmap.len() {
return Ok(0); }
let available = mmap.len() - start;
let to_read = buf.len().min(available);
buf[..to_read].copy_from_slice(&mmap[start..start + to_read]);
Ok(to_read)
}
Some(Inner::Fallback(writer)) => writer.read_at(offset, buf).await,
None => Err(Aria2Error::Io("writer not open".into())),
}
}
async fn truncate(&mut self, length: u64) -> Result<()> {
self.open().await?;
match self.inner.as_mut() {
Some(Inner::Mmap { mmap, .. }) => {
if let Err(e) = mmap.flush_async() {
warn!("mmap flush_async before truncate failed: {}", e);
}
self.inner = None;
let mut writer = PositionedDiskWriter::new(&self.path, self.total_size);
writer.open().await?;
writer.truncate(length).await?;
self.inner = Some(Inner::Fallback(writer));
debug!(
"Truncated mmap writer to {} bytes, switched to fallback: {:?}",
length, self.path
);
Ok(())
}
Some(Inner::Fallback(writer)) => writer.truncate(length).await,
None => Err(Aria2Error::Io("writer not open".into())),
}
}
async fn flush(&mut self) -> Result<()> {
match self.inner.as_mut() {
Some(Inner::Mmap { mmap, .. }) => {
mmap.flush_async()
.map_err(|e| Aria2Error::Io(format!("mmap flush_async failed: {}", e)))?;
Ok(())
}
Some(Inner::Fallback(writer)) => writer.flush().await,
None => Ok(()), }
}
async fn len(&self) -> Result<u64> {
match &self.inner {
Some(Inner::Mmap { mmap, .. }) => Ok(mmap.len() as u64),
Some(Inner::Fallback(writer)) => writer.len().await,
None => {
if let Some(size) = self.total_size {
Ok(size)
} else {
Ok(0)
}
}
}
}
fn path(&self) -> &Path {
&self.path
}
async fn close(&mut self) -> Result<()> {
self.flush().await?;
self.inner = None;
self.opened = false;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
#[tokio::test]
async fn test_mmap_writer_basic() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_mmap_basic.bin");
let mut writer = MmapDiskWriter::new(&path, Some(1024));
writer.open().await.unwrap();
writer.write_at(0, b"hello mmap").await.unwrap();
writer.flush().await.unwrap();
let mut buf = [0u8; 10];
let n = writer.read_at(0, &mut buf).await.unwrap();
assert_eq!(n, 10);
assert_eq!(&buf, b"hello mmap");
}
#[tokio::test]
async fn test_mmap_writer_write_at_offset() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_mmap_offset.bin");
let mut writer = MmapDiskWriter::new(&path, Some(512));
writer.open().await.unwrap();
writer.write_at(100, b"offset data").await.unwrap();
writer.flush().await.unwrap();
let mut buf = [0u8; 11];
let n = writer.read_at(100, &mut buf).await.unwrap();
assert_eq!(n, 11);
assert_eq!(&buf, b"offset data");
let mut buf0 = [0xFFu8; 16];
let n0 = writer.read_at(0, &mut buf0).await.unwrap();
assert_eq!(n0, 16, "should read full 16 bytes from zero-filled region");
assert!(
buf0.iter().all(|&b| b == 0),
"offset 0 should be zero-filled, got {:?}",
buf0
);
}
#[tokio::test]
async fn test_mmap_writer_concurrent_writes() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_mmap_concurrent.bin");
let chunk_size: usize = 64 * 1024;
let num_tasks: usize = 4;
let total_size = (chunk_size * num_tasks) as u64;
let mut writer = MmapDiskWriter::new(&path, Some(total_size));
writer.open().await.unwrap();
let writer = Arc::new(tokio::sync::Mutex::new(writer));
let mut handles = Vec::with_capacity(num_tasks);
for i in 0..num_tasks {
let offset = (i as u64) * chunk_size as u64;
let fill = (i as u8) + 1;
let data = bytes::Bytes::from(vec![fill; chunk_size]);
let w = writer.clone();
handles.push(tokio::spawn(async move {
let mut guard = w.lock().await;
guard.write_bytes_at(offset, data).await.unwrap();
}));
}
for handle in handles {
handle.await.unwrap();
}
{
let mut guard = writer.lock().await;
guard.flush().await.unwrap();
}
for i in 0..num_tasks {
let offset = (i as u64) * chunk_size as u64;
let expected = (i as u8) + 1;
let mut buf = vec![0u8; chunk_size];
let mut guard = writer.lock().await;
let n = guard.read_at(offset, &mut buf).await.unwrap();
assert_eq!(n, chunk_size, "read length mismatch at chunk {}", i);
assert!(
buf.iter().all(|&b| b == expected),
"data mismatch in chunk {}",
i
);
}
}
#[tokio::test]
async fn test_mmap_writer_fallback_on_open_failure() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_mmap_fallback.bin");
let mut writer = MmapDiskWriter::new(&path, None);
writer.open().await.unwrap();
writer.write_at(0, b"fallback works").await.unwrap();
writer.flush().await.unwrap();
let mut buf = [0u8; 14];
let n = writer.read_at(0, &mut buf).await.unwrap();
assert_eq!(n, 14);
assert_eq!(&buf, b"fallback works");
writer.truncate(7).await.unwrap();
let len = writer.len().await.unwrap();
assert_eq!(len, 7);
}
#[tokio::test]
async fn test_mmap_writer_truncate_and_len() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_mmap_trunc.bin");
let mut writer = MmapDiskWriter::new(&path, Some(2048));
writer.open().await.unwrap();
let len = writer.len().await.unwrap();
assert_eq!(len, 2048);
writer.write_at(0, b"before truncate").await.unwrap();
writer.flush().await.unwrap();
writer.truncate(512).await.unwrap();
let len = writer.len().await.unwrap();
assert_eq!(len, 512);
let mut buf = [0u8; 15];
let n = writer.read_at(0, &mut buf).await.unwrap();
assert_eq!(n, 15);
assert_eq!(&buf, b"before truncate");
}
#[tokio::test]
async fn test_mmap_writer_write_bytes_at() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_mmap_bytes.bin");
let mut writer = MmapDiskWriter::new(&path, Some(256));
writer.open().await.unwrap();
let data = bytes::Bytes::from(vec![0xAB; 128]);
writer.write_bytes_at(0, data).await.unwrap();
writer.flush().await.unwrap();
let mut buf = [0u8; 128];
let n = writer.read_at(0, &mut buf).await.unwrap();
assert_eq!(n, 128);
assert!(buf.iter().all(|&b| b == 0xAB));
}
#[tokio::test]
async fn test_mmap_writer_close_and_reopen() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_mmap_close.bin");
let mut writer = MmapDiskWriter::new(&path, Some(1024));
writer.open().await.unwrap();
writer.write_at(0, b"before close").await.unwrap();
writer.close().await.unwrap();
assert!(!writer.opened);
writer.open().await.unwrap();
writer.write_at(12, b" after reopen").await.unwrap();
writer.close().await.unwrap();
let content = std::fs::read(&path).unwrap();
assert_eq!(&content[..25], b"before close after reopen");
}
#[tokio::test]
async fn test_mmap_writer_len_before_open() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_mmap_len_before.bin");
let writer = MmapDiskWriter::new(&path, Some(9999));
let len = writer.len().await.unwrap();
assert_eq!(len, 9999, "should return total_size before open");
}
#[tokio::test]
async fn test_mmap_writer_resume_does_not_truncate() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_mmap_resume.bin");
{
let mut w = MmapDiskWriter::new(&path, Some(1024));
w.open().await.unwrap();
w.write_at(0, b"resume-data").await.unwrap();
w.flush().await.unwrap();
}
{
let mut w = MmapDiskWriter::new(&path, Some(1024));
w.open().await.unwrap();
let mut buf = [0u8; 11];
let n = w.read_at(0, &mut buf).await.unwrap();
assert_eq!(n, 11);
assert_eq!(&buf, b"resume-data", "existing data must survive reopen");
}
}
}