use std::path::{Path, PathBuf};
use std::sync::Mutex;
use async_trait::async_trait;
use tracing::debug;
use crate::error::{Aria2Error, Result};
use super::disk_writer::SeekableDiskWriter;
pub struct PositionedDiskWriter {
file: Mutex<Option<std::fs::File>>,
path: PathBuf,
total_size: Option<u64>,
}
impl PositionedDiskWriter {
pub fn new(path: &Path, total_size: Option<u64>) -> Self {
Self {
file: Mutex::new(None),
path: path.to_path_buf(),
total_size,
}
}
pub fn total_size(&self) -> Option<u64> {
self.total_size
}
fn ensure_open_sync(&self) -> Result<()> {
let mut guard = self
.file
.lock()
.map_err(|e| Aria2Error::Io(format!("file mutex poisoned: {e}")))?;
if guard.is_some() {
return Ok(());
}
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);
}
let mut opts = std::fs::OpenOptions::new();
opts.create(true).write(true).read(true);
let file = opts.open(&self.path)?;
debug!("Opened file for positioned I/O: {:?}", 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);
}
}
*guard = Some(file);
Ok(())
}
fn lock_file(&self) -> Result<std::sync::MutexGuard<'_, Option<std::fs::File>>> {
self.file
.lock()
.map_err(|e| Aria2Error::Io(format!("file mutex poisoned: {e}")))
}
#[cfg(unix)]
pub fn raw_fd(&self) -> Option<std::os::unix::io::RawFd> {
let guard = self.file.lock().ok()?;
guard.as_ref().map(std::os::unix::io::AsRawFd::as_raw_fd)
}
}
#[async_trait]
impl SeekableDiskWriter for PositionedDiskWriter {
async fn open(&mut self) -> Result<()> {
self.ensure_open_sync()
}
async fn write_at(&mut self, offset: u64, data: &[u8]) -> Result<()> {
self.ensure_open_sync()?;
let guard = self.lock_file()?;
let file = guard.as_ref().ok_or_else(|| {
Aria2Error::Io("file not open after ensure_open_sync — invariant violated".into())
})?;
write_all_at(file, data, offset)
}
async fn write_bytes_at(&mut self, offset: u64, data: bytes::Bytes) -> Result<()> {
self.write_at(offset, &data).await
}
async fn read_at(&mut self, offset: u64, buf: &mut [u8]) -> Result<usize> {
self.ensure_open_sync()?;
let guard = self.lock_file()?;
let file = guard.as_ref().ok_or_else(|| {
Aria2Error::Io("file not open after ensure_open_sync — invariant violated".into())
})?;
read_exact_at(file, buf, offset)
}
async fn truncate(&mut self, length: u64) -> Result<()> {
self.ensure_open_sync()?;
let guard = self.lock_file()?;
if let Some(ref file) = *guard {
file.set_len(length)?;
}
Ok(())
}
async fn flush(&mut self) -> Result<()> {
let guard = self.lock_file()?;
if let Some(ref file) = *guard {
file.sync_all()?;
}
Ok(())
}
async fn len(&self) -> Result<u64> {
let guard = self.lock_file()?;
if let Some(ref file) = *guard {
Ok(file.metadata()?.len())
} else if let Some(size) = self.total_size {
Ok(size)
} else {
Ok(0)
}
}
fn path(&self) -> &Path {
&self.path
}
}
fn write_all_at(file: &std::fs::File, mut buf: &[u8], mut offset: u64) -> Result<()> {
while !buf.is_empty() {
let n = positioned_write(file, buf, offset)?;
if n == 0 {
return Err(Aria2Error::Io(
"positioned write returned 0 — failed to write whole buffer".into(),
));
}
offset += n as u64;
buf = &buf[n..];
}
Ok(())
}
fn read_exact_at(file: &std::fs::File, buf: &mut [u8], offset: u64) -> Result<usize> {
let mut filled = 0usize;
let mut current_offset = offset;
while filled < buf.len() {
let n = positioned_read(file, &mut buf[filled..], current_offset)?;
if n == 0 {
break; }
filled += n;
current_offset += n as u64;
}
Ok(filled)
}
#[cfg(unix)]
fn positioned_write(file: &std::fs::File, buf: &[u8], offset: u64) -> Result<usize> {
use std::os::unix::fs::FileExt;
Ok(file.write_at(buf, offset)?)
}
#[cfg(unix)]
fn positioned_read(file: &std::fs::File, buf: &mut [u8], offset: u64) -> Result<usize> {
use std::os::unix::fs::FileExt;
Ok(file.read_at(buf, offset)?)
}
#[cfg(windows)]
fn positioned_write(file: &std::fs::File, buf: &[u8], offset: u64) -> Result<usize> {
use std::os::windows::fs::FileExt;
Ok(file.seek_write(buf, offset)?)
}
#[cfg(windows)]
fn positioned_read(file: &std::fs::File, buf: &mut [u8], offset: u64) -> Result<usize> {
use std::os::windows::fs::FileExt;
Ok(file.seek_read(buf, offset)?)
}
#[cfg(not(any(unix, windows)))]
fn positioned_write(_file: &std::fs::File, _buf: &[u8], _offset: u64) -> Result<usize> {
Err(Aria2Error::Io(
"positioned write not supported on this platform".into(),
))
}
#[cfg(not(any(unix, windows)))]
fn positioned_read(_file: &std::fs::File, _buf: &mut [u8], _offset: u64) -> Result<usize> {
Err(Aria2Error::Io(
"positioned read not supported on this platform".into(),
))
}
pub fn create_positioned_writer(
path: &Path,
total_size: Option<u64>,
) -> Box<dyn SeekableDiskWriter> {
#[cfg(all(target_os = "linux", feature = "io_uring"))]
{
Box::new(IoUringDiskWriter::new(path, total_size))
}
#[cfg(not(all(target_os = "linux", feature = "io_uring")))]
{
Box::new(PositionedDiskWriter::new(path, total_size))
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
#[tokio::test]
async fn test_positioned_write_basic() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_basic.bin");
let mut writer = PositionedDiskWriter::new(&path, Some(1024));
writer.open().await.unwrap();
writer.write_at(0, b"hello world").await.unwrap();
writer.flush().await.unwrap();
let mut buf = [0u8; 11];
let n = writer.read_at(0, &mut buf).await.unwrap();
assert_eq!(n, 11);
assert_eq!(&buf, b"hello world");
}
#[tokio::test]
async fn test_positioned_write_at_offset() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_offset.bin");
let mut writer = PositionedDiskWriter::new(&path, None);
writer.open().await.unwrap();
writer.write_at(100, b"data at 100").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"data at 100");
let mut buf0 = [0xFFu8; 12];
let n0 = writer.read_at(0, &mut buf0).await.unwrap();
assert_eq!(n0, 12, "should read full 12 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_positioned_writer_truncate_and_len() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_trunc.bin");
let mut writer = PositionedDiskWriter::new(&path, Some(2048));
writer.open().await.unwrap();
let len = writer.len().await.unwrap();
assert_eq!(len, 2048);
writer.truncate(512).await.unwrap();
let len = writer.len().await.unwrap();
assert_eq!(len, 512);
}
#[tokio::test]
async fn test_positioned_writer_len_before_open() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_len_before_open.bin");
let writer = PositionedDiskWriter::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_positioned_writer_len_no_total_size_before_open() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_len_none.bin");
let writer = PositionedDiskWriter::new(&path, None);
let len = writer.len().await.unwrap();
assert_eq!(len, 0, "should return 0 before open when no total_size");
}
#[tokio::test]
async fn test_positioned_writer_resume_does_not_truncate() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_resume.bin");
{
let mut w = PositionedDiskWriter::new(&path, Some(1024));
w.open().await.unwrap();
w.write_at(0, b"resume-data").await.unwrap();
w.flush().await.unwrap();
}
{
let mut w = PositionedDiskWriter::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");
}
}
#[tokio::test]
async fn test_positioned_writer_creates_parent_dirs() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("nested").join("deep").join("file.bin");
let mut writer = PositionedDiskWriter::new(&path, Some(64));
writer.open().await.unwrap();
writer.write_at(0, b"x").await.unwrap();
writer.flush().await.unwrap();
assert!(path.exists(), "file should be created with parent dirs");
}
#[tokio::test]
async fn test_concurrent_writes_non_overlapping() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_concurrent_shared.bin");
let chunk_size: usize = 64 * 1024;
let num_tasks: usize = 4;
let mut writer = PositionedDiskWriter::new(&path, Some((chunk_size * num_tasks) as u64));
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();
}
let content = tokio::fs::read(&path).await.unwrap();
assert_eq!(content.len(), chunk_size * num_tasks);
for i in 0..num_tasks {
let start = i * chunk_size;
let expected = (i as u8) + 1;
let chunk = &content[start..start + chunk_size];
assert!(
chunk.iter().all(|&b| b == expected),
"data mismatch in task {} chunk",
i
);
}
}
#[tokio::test]
async fn test_concurrent_writes_separate_writers() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_concurrent_sep.bin");
let chunk_size: usize = 64 * 1024;
let num_tasks: usize = 4;
{
let mut w0 = PositionedDiskWriter::new(&path, Some((chunk_size * num_tasks) as u64));
w0.open().await.unwrap();
w0.flush().await.unwrap();
}
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 = vec![fill; chunk_size];
let path_clone = path.clone();
handles.push(tokio::spawn(async move {
let mut w = PositionedDiskWriter::new(&path_clone, None);
w.open().await.unwrap();
w.write_at(offset, &data).await.unwrap();
w.flush().await.unwrap();
}));
}
for handle in handles {
handle.await.unwrap();
}
let content = tokio::fs::read(&path).await.unwrap();
assert_eq!(content.len(), chunk_size * num_tasks);
for i in 0..num_tasks {
let start = i * chunk_size;
let expected = (i as u8) + 1;
let chunk = &content[start..start + chunk_size];
assert!(
chunk.iter().all(|&b| b == expected),
"data mismatch in separate-writer task {} chunk",
i
);
}
}
#[tokio::test]
async fn test_positioned_writer_write_bytes_at_zero_copy() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("test_zero_copy.bin");
let mut writer = PositionedDiskWriter::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));
}
}
#[cfg(all(target_os = "linux", feature = "io_uring"))]
mod io_uring_backend {
use std::path::{Path, PathBuf};
use async_trait::async_trait;
use tracing::debug;
use crate::error::{Aria2Error, Result};
use super::SeekableDiskWriter;
pub struct IoUringDiskWriter {
file: Option<tokio_uring::fs::File>,
path: PathBuf,
total_size: Option<u64>,
}
impl IoUringDiskWriter {
pub fn new(path: &Path, total_size: Option<u64>) -> Self {
Self {
file: None,
path: path.to_path_buf(),
total_size,
}
}
pub fn total_size(&self) -> Option<u64> {
self.total_size
}
}
#[async_trait]
impl SeekableDiskWriter for IoUringDiskWriter {
async fn open(&mut self) -> Result<()> {
if self.file.is_some() {
return Ok(());
}
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);
}
if let Some(size) = self.total_size {
let std_file = std::fs::OpenOptions::new()
.create(true)
.write(true)
.read(true)
.open(&self.path)?;
let current_size = std_file.metadata()?.len();
if current_size == 0 && size > 0 {
std_file.set_len(size)?;
debug!("Pre-allocated file to {} bytes: {:?}", size, self.path);
}
drop(std_file);
}
let file = tokio_uring::fs::OpenOptions::new()
.create(true)
.write(true)
.read(true)
.open(&self.path)
.await
.map_err(|e| Aria2Error::Io(format!("io_uring open failed: {e}")))?;
debug!("Opened file for io_uring positioned I/O: {:?}", self.path);
self.file = Some(file);
Ok(())
}
async fn write_at(&mut self, offset: u64, data: &[u8]) -> Result<()> {
let file = self
.file
.as_ref()
.ok_or_else(|| Aria2Error::Io("io_uring file not open".into()))?;
write_all_at_uring(file, data, offset).await
}
async fn write_bytes_at(&mut self, offset: u64, data: bytes::Bytes) -> Result<()> {
self.write_at(offset, &data).await
}
async fn read_at(&mut self, offset: u64, buf: &mut [u8]) -> Result<usize> {
let file = self
.file
.as_ref()
.ok_or_else(|| Aria2Error::Io("io_uring file not open".into()))?;
read_exact_at_uring(file, buf, offset).await
}
async fn truncate(&mut self, length: u64) -> Result<()> {
if let Some(file) = self.file.take() {
let _ = file.close().await.map_err(|e| {
Aria2Error::Io(format!("io_uring close during truncate failed: {e}"))
})?;
}
std::fs::OpenOptions::new()
.write(true)
.open(&self.path)?
.set_len(length)?;
let file = tokio_uring::fs::OpenOptions::new()
.create(true)
.write(true)
.read(true)
.open(&self.path)
.await
.map_err(|e| {
Aria2Error::Io(format!("io_uring reopen after truncate failed: {e}"))
})?;
self.file = Some(file);
Ok(())
}
async fn flush(&mut self) -> Result<()> {
if let Some(ref file) = self.file {
file.sync_all()
.await
.map_err(|e| Aria2Error::Io(format!("io_uring sync_all failed: {e}")))?;
}
Ok(())
}
async fn len(&self) -> Result<u64> {
if self.file.is_some() {
Ok(std::fs::metadata(&self.path)
.map(|m| m.len())
.unwrap_or_else(|_| self.total_size.unwrap_or(0)))
} else if let Some(size) = self.total_size {
Ok(size)
} else {
Ok(0)
}
}
fn path(&self) -> &Path {
&self.path
}
async fn close(&mut self) -> Result<()> {
if let Some(file) = self.file.take() {
file.close()
.await
.map_err(|e| Aria2Error::Io(format!("io_uring close failed: {e}")))?;
}
Ok(())
}
}
async fn write_all_at_uring(
file: &tokio_uring::fs::File,
mut buf: &[u8],
mut offset: u64,
) -> Result<()> {
while !buf.is_empty() {
let (res, _) = file.write_at(buf, offset).await;
let n = res.map_err(|e| Aria2Error::Io(format!("io_uring write_at failed: {e}")))?;
if n == 0 {
return Err(Aria2Error::Io(
"io_uring write_at returned 0 — failed to write whole buffer".into(),
));
}
offset += n as u64;
buf = &buf[n..];
}
Ok(())
}
async fn read_exact_at_uring(
file: &tokio_uring::fs::File,
mut buf: &mut [u8],
mut offset: u64,
) -> Result<usize> {
let mut filled = 0usize;
while !buf.is_empty() {
let (res, returned_buf) = file.read_at(buf, offset).await;
let n = res.map_err(|e| Aria2Error::Io(format!("io_uring read_at failed: {e}")))?;
if n == 0 {
break; }
filled += n;
offset += n as u64;
buf = &mut returned_buf[n..];
}
Ok(filled)
}
}
#[cfg(all(target_os = "linux", feature = "io_uring"))]
pub use io_uring_backend::IoUringDiskWriter;
#[cfg(all(test, target_os = "linux", feature = "io_uring"))]
mod io_uring_tests {
use super::IoUringDiskWriter;
#[test]
fn test_iouring_basic_write_read() {
tokio_uring::start(async {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("iouring_basic.bin");
let mut writer = IoUringDiskWriter::new(&path, Some(1024));
writer.open().await.unwrap();
writer.write_at(0, b"hello io_uring").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"hello io_uring");
});
}
#[test]
fn test_iouring_write_at_offset() {
tokio_uring::start(async {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("iouring_offset.bin");
let mut writer = IoUringDiskWriter::new(&path, None);
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");
});
}
#[test]
fn test_iouring_truncate_and_len() {
tokio_uring::start(async {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("iouring_trunc.bin");
let mut writer = IoUringDiskWriter::new(&path, Some(2048));
writer.open().await.unwrap();
let len = writer.len().await.unwrap();
assert_eq!(len, 2048);
writer.truncate(512).await.unwrap();
let len = writer.len().await.unwrap();
assert_eq!(len, 512);
});
}
#[test]
fn test_iouring_close_reopen() {
tokio_uring::start(async {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("iouring_close.bin");
let mut writer = IoUringDiskWriter::new(&path, None);
writer.open().await.unwrap();
writer.write_at(0, b"before close").await.unwrap();
writer.flush().await.unwrap();
writer.close().await.unwrap();
writer.open().await.unwrap();
writer.write_at(12, b" after reopen").await.unwrap();
writer.flush().await.unwrap();
writer.close().await.unwrap();
let content = std::fs::read(&path).unwrap();
assert_eq!(&content, b"before close after reopen");
});
}
#[test]
fn test_iouring_resume_does_not_truncate() {
tokio_uring::start(async {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("iouring_resume.bin");
{
let mut w = IoUringDiskWriter::new(&path, Some(1024));
w.open().await.unwrap();
w.write_at(0, b"resume-data").await.unwrap();
w.flush().await.unwrap();
w.close().await.unwrap();
}
{
let mut w = IoUringDiskWriter::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");
}
});
}
#[test]
fn test_iouring_creates_parent_dirs() {
tokio_uring::start(async {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("nested").join("deep").join("file.bin");
let mut writer = IoUringDiskWriter::new(&path, Some(64));
writer.open().await.unwrap();
writer.write_at(0, b"x").await.unwrap();
writer.flush().await.unwrap();
assert!(path.exists(), "file should be created with parent dirs");
});
}
#[test]
fn test_iouring_write_bytes_at() {
tokio_uring::start(async {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("iouring_bytes.bin");
let mut writer = IoUringDiskWriter::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));
});
}
}