revault_lockbox_api 0.0.1

reVault lockbox API to create and manage lockboxes
Documentation
use crate::{Error, Result};
use std::cell::RefCell;
use std::collections::BTreeMap;
use std::ffi::OsString;
use std::fs::{self, File, OpenOptions};
use std::io::{Read, Seek, SeekFrom, Write};
use std::path::{Path, PathBuf};
use std::thread;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};

const DEFAULT_LOCK_TIMEOUT: Duration = Duration::from_secs(30);
const LOCK_POLL_INTERVAL: Duration = Duration::from_millis(100);

#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FileLockScope {
    Lockbox,
    Vault,
}

impl FileLockScope {
    fn as_str(self) -> &'static str {
        match self {
            Self::Lockbox => "lockbox",
            Self::Vault => "vault",
        }
    }
}

#[derive(Debug)]
pub struct ScopedFileLock {
    lock_path: PathBuf,
    #[cfg(unix)]
    file: Option<File>,
    #[cfg(not(unix))]
    owns_lock_file: bool,
}

impl ScopedFileLock {
    pub fn acquire(target: &Path, scope: FileLockScope) -> Result<Self> {
        let lock_path = lock_path_for(target);
        if enter_thread_lock(&lock_path) {
            return Ok(Self {
                lock_path,
                #[cfg(unix)]
                file: None,
                #[cfg(not(unix))]
                owns_lock_file: false,
            });
        }
        let timeout = lock_timeout();
        let started = Instant::now();
        loop {
            match try_acquire(target, &lock_path, scope) {
                Ok(lock) => return Ok(lock),
                Err(AcquireFailure::Busy(owner)) => {
                    if started.elapsed() >= timeout {
                        leave_thread_lock(&lock_path);
                        return Err(timeout_error(target, scope, timeout, owner.as_deref()));
                    }
                    thread::sleep(LOCK_POLL_INTERVAL);
                }
                Err(AcquireFailure::Io(err)) => {
                    leave_thread_lock(&lock_path);
                    return Err(Error::Io(err));
                }
            }
        }
    }
}

#[cfg(unix)]
impl Drop for ScopedFileLock {
    fn drop(&mut self) {
        if !leave_thread_lock(&self.lock_path) {
            return;
        }
        if let Some(file) = &self.file {
            use std::os::fd::AsRawFd;

            // SAFETY: this releases the same valid descriptor locked in
            // `try_acquire`.
            let _ = unsafe { libc::flock(file.as_raw_fd(), libc::LOCK_UN) };
        }
    }
}

#[cfg(not(unix))]
impl Drop for ScopedFileLock {
    fn drop(&mut self) {
        if leave_thread_lock(&self.lock_path) && self.owns_lock_file {
            let _ = fs::remove_file(&self.lock_path);
        }
    }
}

enum AcquireFailure {
    Busy(Option<String>),
    Io(String),
}

#[cfg(unix)]
fn try_acquire(
    target: &Path,
    lock_path: &Path,
    scope: FileLockScope,
) -> std::result::Result<ScopedFileLock, AcquireFailure> {
    use std::os::fd::AsRawFd;

    if let Some(parent) = lock_path.parent() {
        fs::create_dir_all(parent).map_err(|err| AcquireFailure::Io(err.to_string()))?;
    }
    let mut file = OpenOptions::new()
        .read(true)
        .write(true)
        .create(true)
        .truncate(false)
        .open(lock_path)
        .map_err(|err| AcquireFailure::Io(format!("open {}: {err}", lock_path.display())))?;
    // SAFETY: flock operates on a valid file descriptor owned by `file`.
    let rc = unsafe { libc::flock(file.as_raw_fd(), libc::LOCK_EX | libc::LOCK_NB) };
    if rc == 0 {
        write_owner_metadata(&mut file, target, scope)
            .map_err(|err| AcquireFailure::Io(format!("write {}: {err}", lock_path.display())))?;
        return Ok(ScopedFileLock {
            lock_path: lock_path.to_path_buf(),
            file: Some(file),
        });
    }
    let err = std::io::Error::last_os_error();
    if err.raw_os_error() == Some(libc::EWOULDBLOCK) || err.raw_os_error() == Some(libc::EAGAIN) {
        return Err(AcquireFailure::Busy(read_owner_metadata(lock_path)));
    }
    Err(AcquireFailure::Io(err.to_string()))
}

#[cfg(not(unix))]
fn try_acquire(
    target: &Path,
    lock_path: &Path,
    scope: FileLockScope,
) -> std::result::Result<ScopedFileLock, AcquireFailure> {
    if let Some(parent) = lock_path.parent() {
        fs::create_dir_all(parent).map_err(|err| AcquireFailure::Io(err.to_string()))?;
    }
    let mut file = match OpenOptions::new()
        .write(true)
        .create_new(true)
        .open(&lock_path)
    {
        Ok(file) => file,
        Err(err) if err.kind() == std::io::ErrorKind::AlreadyExists => {
            if lock_file_is_stale(&lock_path) {
                let _ = fs::remove_file(&lock_path);
                return Err(AcquireFailure::Busy(None));
            }
            return Err(AcquireFailure::Busy(read_owner_metadata(&lock_path)));
        }
        Err(err) => {
            return Err(AcquireFailure::Io(format!(
                "create {}: {err}",
                lock_path.display()
            )));
        }
    };
    write_owner_metadata(&mut file, target, scope)
        .map_err(|err| AcquireFailure::Io(format!("write {}: {err}", lock_path.display())))?;
    Ok(ScopedFileLock {
        lock_path: lock_path.to_path_buf(),
        owns_lock_file: true,
    })
}

thread_local! {
    static THREAD_LOCKS: RefCell<BTreeMap<PathBuf, usize>> = const {
        RefCell::new(BTreeMap::new())
    };
}

fn enter_thread_lock(lock_path: &Path) -> bool {
    THREAD_LOCKS.with(|locks| {
        let mut locks = locks.borrow_mut();
        let count = locks.entry(lock_path.to_path_buf()).or_insert(0);
        let nested = *count > 0;
        *count = count.saturating_add(1);
        nested
    })
}

fn leave_thread_lock(lock_path: &Path) -> bool {
    THREAD_LOCKS.with(|locks| {
        let mut locks = locks.borrow_mut();
        let Some(count) = locks.get_mut(lock_path) else {
            return true;
        };
        *count = count.saturating_sub(1);
        if *count == 0 {
            locks.remove(lock_path);
            true
        } else {
            false
        }
    })
}

/// Returns the hidden sidecar path used to coordinate writes to `target`.
pub fn lock_path_for(target: &Path) -> PathBuf {
    let Some(file_name) = target.file_name() else {
        let mut path = OsString::from(target.as_os_str());
        path.push(".lock");
        return PathBuf::from(path);
    };
    let mut lock_name = OsString::from(".");
    lock_name.push(file_name);
    lock_name.push(".lock");
    target.with_file_name(lock_name)
}

fn lock_timeout() -> Duration {
    std::env::var("LOCKBOX_LOCK_TIMEOUT_MS")
        .ok()
        .and_then(|value| value.parse::<u64>().ok())
        .map(Duration::from_millis)
        .unwrap_or(DEFAULT_LOCK_TIMEOUT)
}

fn write_owner_metadata(
    file: &mut File,
    target: &Path,
    scope: FileLockScope,
) -> std::io::Result<()> {
    file.set_len(0)?;
    file.seek(SeekFrom::Start(0))?;
    let now_ms = SystemTime::now()
        .duration_since(UNIX_EPOCH)
        .unwrap_or_default()
        .as_millis();
    let exe = std::env::current_exe()
        .ok()
        .map(|path| path.display().to_string())
        .unwrap_or_else(|| "unknown".to_string());
    let user = std::env::var("USER")
        .or_else(|_| std::env::var("USERNAME"))
        .unwrap_or_else(|_| "unknown".to_string());
    writeln!(file, "scope={}", scope.as_str())?;
    writeln!(file, "target={}", target.display())?;
    writeln!(file, "pid={}", std::process::id())?;
    writeln!(file, "user={user}")?;
    writeln!(file, "exe={exe}")?;
    writeln!(file, "created_unix_ms={now_ms}")?;
    file.sync_data()
}

fn read_owner_metadata(path: &Path) -> Option<String> {
    let mut text = String::new();
    File::open(path).ok()?.read_to_string(&mut text).ok()?;
    let pid = metadata_value(&text, "pid").unwrap_or("unknown");
    let created = metadata_value(&text, "created_unix_ms")
        .and_then(|value| value.parse::<u128>().ok())
        .map(format_unix_millis_datetime)
        .or_else(|| metadata_value(&text, "created_unix_ms").map(str::to_string))
        .unwrap_or_else(|| "unknown".to_string());
    Some(format!(" by pid {pid} since {created}"))
}

fn timeout_error(
    target: &Path,
    scope: FileLockScope,
    timeout: Duration,
    owner: Option<&str>,
) -> Error {
    Error::LockUnavailable(format!(
        "{} {} is locked{}; timed out after {}s",
        scope.as_str(),
        target.display(),
        owner.unwrap_or(""),
        timeout.as_secs()
    ))
}

fn metadata_value<'a>(text: &'a str, key: &str) -> Option<&'a str> {
    text.lines()
        .find_map(|line| line.strip_prefix(key)?.strip_prefix('='))
}

fn format_unix_millis_datetime(unix_ms: u128) -> String {
    let rounded_seconds = ((unix_ms + 500) / 1000).min(i64::MAX as u128) as i64;
    let days = rounded_seconds.div_euclid(86_400);
    let seconds_of_day = rounded_seconds.rem_euclid(86_400);
    let (year, month, day) = civil_from_days(days);
    let hour = seconds_of_day / 3_600;
    let minute = (seconds_of_day % 3_600) / 60;
    let second = seconds_of_day % 60;
    format!("{year:04}-{month:02}-{day:02} {hour:02}:{minute:02}:{second:02}")
}

fn civil_from_days(days_since_unix_epoch: i64) -> (i64, u32, u32) {
    let z = days_since_unix_epoch + 719_468;
    let era = if z >= 0 { z } else { z - 146_096 } / 146_097;
    let day_of_era = z - era * 146_097;
    let year_of_era =
        (day_of_era - day_of_era / 1_460 + day_of_era / 36_524 - day_of_era / 146_096) / 365;
    let year = year_of_era + era * 400;
    let day_of_year = day_of_era - (365 * year_of_era + year_of_era / 4 - year_of_era / 100);
    let month_prime = (5 * day_of_year + 2) / 153;
    let day = day_of_year - (153 * month_prime + 2) / 5 + 1;
    let month = month_prime + if month_prime < 10 { 3 } else { -9 };
    let year = year + if month <= 2 { 1 } else { 0 };
    (year, month as u32, day as u32)
}

#[cfg(not(unix))]
fn lock_file_is_stale(path: &Path) -> bool {
    let Ok(text) = fs::read_to_string(path) else {
        return false;
    };
    let Some(pid) = metadata_value(&text, "pid").and_then(|pid| pid.parse::<u32>().ok()) else {
        return false;
    };
    !process_exists(pid)
}

#[cfg(not(unix))]
fn process_exists(pid: u32) -> bool {
    let pid = sysinfo::Pid::from_u32(pid);
    let mut system = sysinfo::System::new();
    system.refresh_processes(sysinfo::ProcessesToUpdate::Some(&[pid]), true);
    system.process(pid).is_some()
}

#[cfg(test)]
mod tests {
    use super::{format_unix_millis_datetime, lock_path_for};
    use std::path::Path;

    #[test]
    fn lock_file_is_hidden_beside_target() {
        assert_eq!(
            lock_path_for(Path::new("/tmp/secrets.lbox")),
            Path::new("/tmp/.secrets.lbox.lock")
        );
        assert_eq!(
            lock_path_for(Path::new("relative.lbox")),
            Path::new(".relative.lbox.lock")
        );
    }

    #[test]
    fn unix_millis_format_is_human_readable_datetime() {
        assert_eq!(format_unix_millis_datetime(0), "1970-01-01 00:00:00");
        assert_eq!(
            format_unix_millis_datetime(1_704_067_201_234),
            "2024-01-01 00:00:01"
        );
        assert_eq!(
            format_unix_millis_datetime(1_704_067_201_500),
            "2024-01-01 00:00:02"
        );
    }
}