1use std::path::{Path, PathBuf};
24
25use async_trait::async_trait;
26use memmap2::MmapMut;
27use tracing::{debug, warn};
28
29use crate::error::{Aria2Error, Result};
30
31use super::disk_writer::SeekableDiskWriter;
32use super::positioned_disk_writer::PositionedDiskWriter;
33
34enum Inner {
36 Mmap {
39 #[allow(dead_code)]
43 file: std::fs::File,
44 mmap: MmapMut,
46 },
47 Fallback(PositionedDiskWriter),
50}
51
52pub struct MmapDiskWriter {
72 inner: Option<Inner>,
74 path: PathBuf,
75 total_size: Option<u64>,
76 opened: bool,
77}
78
79impl MmapDiskWriter {
80 pub fn new(path: &Path, total_size: Option<u64>) -> Self {
86 Self {
87 inner: None,
88 path: path.to_path_buf(),
89 total_size,
90 opened: false,
91 }
92 }
93
94 fn ensure_parent_dirs(&self) -> Result<()> {
96 if let Some(parent) = self.path.parent()
97 && !parent.as_os_str().is_empty()
98 && !parent.exists()
99 {
100 std::fs::create_dir_all(parent)?;
101 debug!("Created parent directories for {:?}", self.path);
102 }
103 Ok(())
104 }
105
106 fn open_sync(&mut self) -> Result<()> {
111 self.ensure_parent_dirs()?;
112
113 let file = std::fs::OpenOptions::new()
114 .create(true)
115 .write(true)
116 .read(true)
117 .truncate(false) .open(&self.path)?;
119
120 if let Some(size) = self.total_size {
122 let current_size = file.metadata()?.len();
123 if current_size == 0 && size > 0 {
124 file.set_len(size)?;
125 debug!("Pre-allocated file to {} bytes: {:?}", size, self.path);
126 }
127 }
128
129 let file_size = file.metadata()?.len();
130 if file_size == 0 {
131 warn!(
133 "File size is 0, cannot mmap: {:?}, using positioned I/O fallback",
134 self.path
135 );
136 drop(file);
140 let writer = PositionedDiskWriter::new(&self.path, self.total_size);
141 self.inner = Some(Inner::Fallback(writer));
142 return Ok(());
143 }
144
145 match unsafe { MmapMut::map_mut(&file) } {
150 Ok(mmap) => {
151 debug!(
152 "Created mmap for {:?}, size: {} bytes",
153 self.path, file_size
154 );
155 self.inner = Some(Inner::Mmap { file, mmap });
156 }
157 Err(e) => {
158 warn!(
159 "mmap failed for {:?}: {}, using positioned I/O fallback",
160 self.path, e
161 );
162 drop(file);
166 let writer = PositionedDiskWriter::new(&self.path, self.total_size);
167 self.inner = Some(Inner::Fallback(writer));
168 }
169 }
170 Ok(())
171 }
172}
173
174#[async_trait]
175impl SeekableDiskWriter for MmapDiskWriter {
176 async fn open(&mut self) -> Result<()> {
177 if self.opened {
178 return Ok(());
179 }
180
181 if self.inner.is_none() {
182 self.open_sync()?;
183 }
184
185 if let Some(Inner::Fallback(ref mut writer)) = self.inner {
187 writer.open().await?;
188 }
189
190 self.opened = true;
191 Ok(())
192 }
193
194 async fn write_at(&mut self, offset: u64, data: &[u8]) -> Result<()> {
195 self.open().await?;
196 match self.inner.as_mut() {
197 Some(Inner::Mmap { mmap, .. }) => {
198 let start = offset as usize;
199 let end = start
200 .checked_add(data.len())
201 .ok_or_else(|| Aria2Error::Io("write offset + length overflow".into()))?;
202 if end > mmap.len() {
203 return Err(Aria2Error::Io(format!(
204 "write at offset {} len {} exceeds mmap size {}",
205 offset,
206 data.len(),
207 mmap.len()
208 )));
209 }
210 mmap[start..end].copy_from_slice(data);
211 Ok(())
212 }
213 Some(Inner::Fallback(writer)) => writer.write_at(offset, data).await,
214 None => Err(Aria2Error::Io("writer not open".into())),
215 }
216 }
217
218 async fn write_bytes_at(&mut self, offset: u64, data: bytes::Bytes) -> Result<()> {
224 self.open().await?;
225 match self.inner.as_mut() {
226 Some(Inner::Mmap { mmap, .. }) => {
227 let start = offset as usize;
228 let end = start
229 .checked_add(data.len())
230 .ok_or_else(|| Aria2Error::Io("write offset + length overflow".into()))?;
231 if end > mmap.len() {
232 return Err(Aria2Error::Io(format!(
233 "write at offset {} len {} exceeds mmap size {}",
234 offset,
235 data.len(),
236 mmap.len()
237 )));
238 }
239 mmap[start..end].copy_from_slice(&data);
240 Ok(())
241 }
242 Some(Inner::Fallback(writer)) => writer.write_bytes_at(offset, data).await,
243 None => Err(Aria2Error::Io("writer not open".into())),
244 }
245 }
246
247 async fn read_at(&mut self, offset: u64, buf: &mut [u8]) -> Result<usize> {
248 self.open().await?;
249 match self.inner.as_mut() {
250 Some(Inner::Mmap { mmap, .. }) => {
251 let start = offset as usize;
252 if start >= mmap.len() {
253 return Ok(0); }
255 let available = mmap.len() - start;
256 let to_read = buf.len().min(available);
257 buf[..to_read].copy_from_slice(&mmap[start..start + to_read]);
258 Ok(to_read)
259 }
260 Some(Inner::Fallback(writer)) => writer.read_at(offset, buf).await,
261 None => Err(Aria2Error::Io("writer not open".into())),
262 }
263 }
264
265 async fn truncate(&mut self, length: u64) -> Result<()> {
266 self.open().await?;
267 match self.inner.as_mut() {
268 Some(Inner::Mmap { mmap, .. }) => {
269 if let Err(e) = mmap.flush_async() {
276 warn!("mmap flush_async before truncate failed: {}", e);
277 }
278 self.inner = None;
281 let mut writer = PositionedDiskWriter::new(&self.path, self.total_size);
282 writer.open().await?;
283 writer.truncate(length).await?;
284 self.inner = Some(Inner::Fallback(writer));
285 debug!(
286 "Truncated mmap writer to {} bytes, switched to fallback: {:?}",
287 length, self.path
288 );
289 Ok(())
290 }
291 Some(Inner::Fallback(writer)) => writer.truncate(length).await,
292 None => Err(Aria2Error::Io("writer not open".into())),
293 }
294 }
295
296 async fn flush(&mut self) -> Result<()> {
297 match self.inner.as_mut() {
298 Some(Inner::Mmap { mmap, .. }) => {
299 mmap.flush_async()
306 .map_err(|e| Aria2Error::Io(format!("mmap flush_async failed: {}", e)))?;
307 Ok(())
308 }
309 Some(Inner::Fallback(writer)) => writer.flush().await,
310 None => Ok(()), }
312 }
313
314 async fn len(&self) -> Result<u64> {
315 match &self.inner {
316 Some(Inner::Mmap { mmap, .. }) => Ok(mmap.len() as u64),
317 Some(Inner::Fallback(writer)) => writer.len().await,
318 None => {
319 if let Some(size) = self.total_size {
320 Ok(size)
321 } else {
322 Ok(0)
323 }
324 }
325 }
326 }
327
328 fn path(&self) -> &Path {
329 &self.path
330 }
331
332 async fn close(&mut self) -> Result<()> {
336 self.flush().await?;
337 self.inner = None;
339 self.opened = false;
340 Ok(())
341 }
342}
343
344#[cfg(test)]
349mod tests {
350 use super::*;
351 use std::sync::Arc;
352
353 #[tokio::test]
354 async fn test_mmap_writer_basic() {
355 let dir = tempfile::tempdir().unwrap();
356 let path = dir.path().join("test_mmap_basic.bin");
357
358 let mut writer = MmapDiskWriter::new(&path, Some(1024));
359 writer.open().await.unwrap();
360 writer.write_at(0, b"hello mmap").await.unwrap();
361 writer.flush().await.unwrap();
362
363 let mut buf = [0u8; 10];
364 let n = writer.read_at(0, &mut buf).await.unwrap();
365 assert_eq!(n, 10);
366 assert_eq!(&buf, b"hello mmap");
367 }
368
369 #[tokio::test]
370 async fn test_mmap_writer_write_at_offset() {
371 let dir = tempfile::tempdir().unwrap();
372 let path = dir.path().join("test_mmap_offset.bin");
373
374 let mut writer = MmapDiskWriter::new(&path, Some(512));
375 writer.open().await.unwrap();
376
377 writer.write_at(100, b"offset data").await.unwrap();
379 writer.flush().await.unwrap();
380
381 let mut buf = [0u8; 11];
383 let n = writer.read_at(100, &mut buf).await.unwrap();
384 assert_eq!(n, 11);
385 assert_eq!(&buf, b"offset data");
386
387 let mut buf0 = [0xFFu8; 16];
389 let n0 = writer.read_at(0, &mut buf0).await.unwrap();
390 assert_eq!(n0, 16, "should read full 16 bytes from zero-filled region");
391 assert!(
392 buf0.iter().all(|&b| b == 0),
393 "offset 0 should be zero-filled, got {:?}",
394 buf0
395 );
396 }
397
398 #[tokio::test]
399 async fn test_mmap_writer_concurrent_writes() {
400 let dir = tempfile::tempdir().unwrap();
404 let path = dir.path().join("test_mmap_concurrent.bin");
405
406 let chunk_size: usize = 64 * 1024;
407 let num_tasks: usize = 4;
408 let total_size = (chunk_size * num_tasks) as u64;
409
410 let mut writer = MmapDiskWriter::new(&path, Some(total_size));
411 writer.open().await.unwrap();
412 let writer = Arc::new(tokio::sync::Mutex::new(writer));
413
414 let mut handles = Vec::with_capacity(num_tasks);
415 for i in 0..num_tasks {
416 let offset = (i as u64) * chunk_size as u64;
417 let fill = (i as u8) + 1;
418 let data = bytes::Bytes::from(vec![fill; chunk_size]);
419 let w = writer.clone();
420 handles.push(tokio::spawn(async move {
421 let mut guard = w.lock().await;
422 guard.write_bytes_at(offset, data).await.unwrap();
423 }));
424 }
425
426 for handle in handles {
427 handle.await.unwrap();
428 }
429
430 {
431 let mut guard = writer.lock().await;
432 guard.flush().await.unwrap();
433 }
434
435 for i in 0..num_tasks {
437 let offset = (i as u64) * chunk_size as u64;
438 let expected = (i as u8) + 1;
439 let mut buf = vec![0u8; chunk_size];
440 let mut guard = writer.lock().await;
441 let n = guard.read_at(offset, &mut buf).await.unwrap();
442 assert_eq!(n, chunk_size, "read length mismatch at chunk {}", i);
443 assert!(
444 buf.iter().all(|&b| b == expected),
445 "data mismatch in chunk {}",
446 i
447 );
448 }
449 }
450
451 #[tokio::test]
452 async fn test_mmap_writer_fallback_on_open_failure() {
453 let dir = tempfile::tempdir().unwrap();
457 let path = dir.path().join("test_mmap_fallback.bin");
458
459 let mut writer = MmapDiskWriter::new(&path, None);
460 writer.open().await.unwrap();
461
462 writer.write_at(0, b"fallback works").await.unwrap();
464 writer.flush().await.unwrap();
465
466 let mut buf = [0u8; 14];
467 let n = writer.read_at(0, &mut buf).await.unwrap();
468 assert_eq!(n, 14);
469 assert_eq!(&buf, b"fallback works");
470
471 writer.truncate(7).await.unwrap();
474 let len = writer.len().await.unwrap();
475 assert_eq!(len, 7);
476 }
477
478 #[tokio::test]
479 async fn test_mmap_writer_truncate_and_len() {
480 let dir = tempfile::tempdir().unwrap();
481 let path = dir.path().join("test_mmap_trunc.bin");
482
483 let mut writer = MmapDiskWriter::new(&path, Some(2048));
484 writer.open().await.unwrap();
485
486 let len = writer.len().await.unwrap();
488 assert_eq!(len, 2048);
489
490 writer.write_at(0, b"before truncate").await.unwrap();
492 writer.flush().await.unwrap();
493
494 writer.truncate(512).await.unwrap();
496 let len = writer.len().await.unwrap();
497 assert_eq!(len, 512);
498
499 let mut buf = [0u8; 15];
501 let n = writer.read_at(0, &mut buf).await.unwrap();
502 assert_eq!(n, 15);
503 assert_eq!(&buf, b"before truncate");
504 }
505
506 #[tokio::test]
507 async fn test_mmap_writer_write_bytes_at() {
508 let dir = tempfile::tempdir().unwrap();
509 let path = dir.path().join("test_mmap_bytes.bin");
510
511 let mut writer = MmapDiskWriter::new(&path, Some(256));
512 writer.open().await.unwrap();
513
514 let data = bytes::Bytes::from(vec![0xAB; 128]);
515 writer.write_bytes_at(0, data).await.unwrap();
516 writer.flush().await.unwrap();
517
518 let mut buf = [0u8; 128];
519 let n = writer.read_at(0, &mut buf).await.unwrap();
520 assert_eq!(n, 128);
521 assert!(buf.iter().all(|&b| b == 0xAB));
522 }
523
524 #[tokio::test]
525 async fn test_mmap_writer_close_and_reopen() {
526 let dir = tempfile::tempdir().unwrap();
527 let path = dir.path().join("test_mmap_close.bin");
528
529 let mut writer = MmapDiskWriter::new(&path, Some(1024));
530 writer.open().await.unwrap();
531 writer.write_at(0, b"before close").await.unwrap();
532 writer.close().await.unwrap();
533 assert!(!writer.opened);
534
535 writer.open().await.unwrap();
536 writer.write_at(12, b" after reopen").await.unwrap();
537 writer.close().await.unwrap();
538
539 let content = std::fs::read(&path).unwrap();
540 assert_eq!(&content[..25], b"before close after reopen");
541 }
542
543 #[tokio::test]
544 async fn test_mmap_writer_len_before_open() {
545 let dir = tempfile::tempdir().unwrap();
546 let path = dir.path().join("test_mmap_len_before.bin");
547
548 let writer = MmapDiskWriter::new(&path, Some(9999));
549 let len = writer.len().await.unwrap();
550 assert_eq!(len, 9999, "should return total_size before open");
551 }
552
553 #[tokio::test]
554 async fn test_mmap_writer_resume_does_not_truncate() {
555 let dir = tempfile::tempdir().unwrap();
558 let path = dir.path().join("test_mmap_resume.bin");
559
560 {
562 let mut w = MmapDiskWriter::new(&path, Some(1024));
563 w.open().await.unwrap();
564 w.write_at(0, b"resume-data").await.unwrap();
565 w.flush().await.unwrap();
566 }
567
568 {
570 let mut w = MmapDiskWriter::new(&path, Some(1024));
571 w.open().await.unwrap();
572 let mut buf = [0u8; 11];
573 let n = w.read_at(0, &mut buf).await.unwrap();
574 assert_eq!(n, 11);
575 assert_eq!(&buf, b"resume-data", "existing data must survive reopen");
576 }
577 }
578}