palace 0.1.0

A tool for mounting datasets into memory for fast loading in deep learning tasks.
Documentation
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)?;

    // Allocate memory by mounting tmpfs
    let tmpfs_size = (source_size as f64 * 1.05) as u64; // 10% extra space
    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,
    )?;

    // Clean up progress file after successful mount
    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(())
}