use crate::utils::{check_memory_availability, check_mount_point, ensure_directory, get_directory_size};
use anyhow::{Context, Result};
use indicatif::{MultiProgress, ProgressBar, ProgressStyle};
use memmap2::MmapOptions;
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
use std::collections::HashSet;
use std::fs::{File, OpenOptions};
use std::path::{Path, PathBuf};
use std::process::Command;
use std::sync::{Arc, Mutex};
#[derive(Serialize, Deserialize, Default)]
struct MountProgress {
completed_files: HashSet<PathBuf>,
}
pub struct MountOptions {
pub source: PathBuf,
pub destination: PathBuf,
pub parallel: usize,
pub use_mmap: bool,
pub compress: bool,
pub memory_threshold: u8,
}
pub fn mount_dataset(options: &MountOptions) -> Result<()> {
check_mount_point(&options.destination)?;
ensure_directory(&options.destination)?;
let source_size = get_directory_size(&options.source)?;
check_memory_availability(source_size, options.memory_threshold)?;
let tmpfs_size = (source_size as f64 * 1.05) as u64; mount_tmpfs(&options.destination, tmpfs_size)?;
let progress_file = options.destination.join(".mount_progress.json");
let mount_progress = load_or_create_progress(&progress_file)?;
if options.use_mmap {
mount_with_mmap(&options.source, &options.destination)
} else {
mount_with_copy(options, mount_progress, &progress_file)
}
}
fn mount_tmpfs(mount_point: &Path, size: u64) -> Result<()> {
let size_mb = size / 1024 / 1024;
let output = Command::new("sudo")
.args(&["mount", "-t", "tmpfs", "-o", &format!("size={}M", size_mb), "tmpfs"])
.arg(mount_point)
.output()
.context("Failed to execute mount command")?;
if !output.status.success() {
let error = String::from_utf8_lossy(&output.stderr);
return Err(anyhow::anyhow!("Failed to mount tmpfs: {}", error));
}
Ok(())
}
fn mount_with_mmap(source: &Path, destination: &Path) -> Result<()> {
let pb = ProgressBar::new_spinner();
pb.set_style(ProgressStyle::default_spinner().template("{spinner:.green} {msg}")?);
pb.set_message("Preparing memory-mapped file...");
let source_size = get_directory_size(source)?;
let file = OpenOptions::new()
.read(true)
.write(true)
.create(true)
.open(destination.join("mmap_file"))?;
file.set_len(source_size)?;
let mut mmap = unsafe { MmapOptions::new().map_mut(&file)? };
pb.set_message("Copying data to memory-mapped file...");
copy_directory_to_mmap(source, &mut mmap, &pb)?;
pb.finish_with_message("Memory-mapped mount completed successfully.");
Ok(())
}
fn copy_directory_to_mmap(source: &Path, mmap: &mut [u8], pb: &ProgressBar) -> Result<()> {
let mut offset = 0;
for entry in std::fs::read_dir(source)? {
let entry = entry?;
let path = entry.path();
if path.is_file() {
let mut file = File::open(&path)?;
let metadata = file.metadata()?;
let len = metadata.len() as usize;
pb.set_message(format!("Copying {}", path.display()));
std::io::copy(&mut file, &mut &mut mmap[offset..offset + len])?;
offset += len;
} else if path.is_dir() {
copy_directory_to_mmap(&path, &mut mmap[offset..], pb)?;
}
}
Ok(())
}
fn load_or_create_progress(progress_file: &Path) -> Result<MountProgress> {
if progress_file.exists() {
let file = File::open(progress_file)?;
Ok(serde_json::from_reader(file)?)
} else {
Ok(MountProgress::default())
}
}
fn save_progress(progress: &MountProgress, progress_file: &Path) -> Result<()> {
let file = File::create(progress_file)?;
serde_json::to_writer(file, progress)?;
Ok(())
}
fn mount_with_copy(options: &MountOptions, progress: MountProgress, progress_file: &Path) -> Result<()> {
let multi_progress = Arc::new(MultiProgress::new());
let style = ProgressStyle::default_bar()
.template("[{elapsed_precise}] {bar:40.cyan/blue} {pos:>7}/{len:7} {msg}")?;
let progress = Arc::new(Mutex::new(progress));
let progress_file = Arc::new(progress_file.to_path_buf());
copy_directory_recursive(
&options.source,
&options.destination,
options,
Arc::clone(&progress),
Arc::clone(&progress_file),
&multi_progress,
&style,
)?;
std::fs::remove_file(&*progress_file)?;
Ok(())
}
fn copy_directory_recursive(
source: &Path,
destination: &Path,
options: &MountOptions,
progress: Arc<Mutex<MountProgress>>,
progress_file: Arc<PathBuf>,
multi_progress: &MultiProgress,
style: &ProgressStyle,
) -> Result<()> {
let entries: Vec<_> = std::fs::read_dir(source)?
.filter_map(Result::ok)
.collect();
entries.into_par_iter()
.with_max_len(options.parallel)
.try_for_each(|entry| {
let path = entry.path();
let dest_path = destination.join(path.strip_prefix(source)?);
let progress_bar = multi_progress.add(ProgressBar::new(0));
progress_bar.set_style(style.clone());
if path.is_file() {
let mut progress_guard = progress.lock().unwrap();
if !progress_guard.completed_files.contains(&dest_path) {
copy_file_with_progress(&path, &dest_path, options.compress, &progress_bar)?;
progress_guard.completed_files.insert(dest_path.clone());
save_progress(&progress_guard, &progress_file)?;
} else {
progress_bar.set_message(format!("Skipping {} (already copied)", path.display()));
progress_bar.finish();
}
} else if path.is_dir() {
std::fs::create_dir_all(&dest_path)?;
copy_directory_recursive(
&path,
&dest_path,
options,
Arc::clone(&progress),
Arc::clone(&progress_file),
multi_progress,
style,
)?;
}
Ok(())
})
}
fn copy_file_with_progress(source: &Path, destination: &Path, compress: bool, progress: &ProgressBar) -> Result<()> {
let mut source_file = File::open(source)?;
let mut dest_file = File::create(destination)?;
let file_size = source_file.metadata()?.len();
progress.set_length(file_size);
progress.set_message(source.file_name().unwrap().to_string_lossy().to_string());
if compress {
let encoder = zstd::Encoder::new(&mut dest_file, 3)?;
std::io::copy(&mut source_file, &mut encoder.auto_finish())?;
} else {
std::io::copy(&mut source_file, &mut dest_file)?;
}
progress.finish_with_message("Done");
Ok(())
}