ztheme 1.6.1

Fast asynchronous Zsh prompt
use std::collections::HashMap;
use std::fs::{self, File, OpenOptions};
use std::io::{self, Read, Write};
use std::os::unix::fs::{MetadataExt as _, PermissionsExt as _};
use std::path::{Path, PathBuf};
use std::sync::Arc;

use super::{
    CACHE_FORMAT_VERSION, CacheKey, Entry, MAX_ENTRIES, MAX_VALUE_BYTES, now_epoch_seconds,
    validate_value,
};

const MAGIC: [u8; 4] = *b"ZTHC";
const MAX_FILE_BYTES: u64 = 9 * 1024 * 1024;
const FUTURE_SKEW_SECONDS: u64 = 5 * 60;

pub(super) fn load(path: &Path) -> io::Result<HashMap<CacheKey, Entry>> {
    validate_private_file(path)?;
    let metadata = path.metadata()?;
    if metadata.len() > MAX_FILE_BYTES {
        return Err(invalid_data("cache file exceeds size limit"));
    }
    let capacity =
        usize::try_from(metadata.len()).map_err(|_| invalid_data("cache file is too large"))?;
    let mut bytes = Vec::with_capacity(capacity);
    File::open(path)?
        .take(MAX_FILE_BYTES.saturating_add(1))
        .read_to_end(&mut bytes)?;
    if u64::try_from(bytes.len()).unwrap_or(u64::MAX) > MAX_FILE_BYTES {
        return Err(invalid_data("cache file exceeds size limit"));
    }
    decode(&bytes)
}

pub(super) fn save(path: &Path, entries: &HashMap<CacheKey, Entry>) -> io::Result<()> {
    let Some(parent) = path.parent() else {
        return Err(invalid_data("cache path has no parent"));
    };
    ensure_private_directory(parent)?;

    let temporary = temporary_path(path);
    let result = (|| {
        let mut file = open_temporary(&temporary)?;
        file.set_permissions(fs::Permissions::from_mode(0o600))?;
        encode(&mut file, entries)?;
        file.flush()?;
        file.sync_all()?;
        fs::rename(&temporary, path)?;
        sync_directory(parent)
    })();

    if result.is_err() {
        let _ = fs::remove_file(&temporary);
    }
    result
}

pub(super) fn clear_all(path: Option<&Path>) -> io::Result<()> {
    let Some(path) = path else {
        return Ok(());
    };
    let Some(directory) = path.parent() else {
        return Ok(());
    };
    let metadata = match fs::symlink_metadata(directory) {
        Ok(metadata) => metadata,
        Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(()),
        Err(error) => return Err(error),
    };
    if !metadata.file_type().is_dir() || metadata.mode() & 0o077 != 0 {
        return Err(invalid_data("cache directory permissions are unsafe"));
    }

    for name in [path.file_name(), temporary_path(path).file_name()]
        .into_iter()
        .flatten()
    {
        let candidate = directory.join(name);
        match fs::remove_file(candidate) {
            Ok(()) => {}
            Err(error) if error.kind() == io::ErrorKind::NotFound => {}
            Err(error) => return Err(error),
        }
    }
    sync_directory(directory)
}

fn decode(bytes: &[u8]) -> io::Result<HashMap<CacheKey, Entry>> {
    let mut decoder = Decoder::new(bytes);
    if decoder.take(MAGIC.len())? != MAGIC {
        return Err(invalid_data("cache file magic is invalid"));
    }
    if decoder.u16()? != CACHE_FORMAT_VERSION {
        return Err(invalid_data("cache file version is unsupported"));
    }
    let count = usize::try_from(decoder.u32()?)
        .map_err(|_| invalid_data("cache entry count is invalid"))?;
    if count > MAX_ENTRIES {
        return Err(invalid_data("cache entry count exceeds limit"));
    }

    let now = now_epoch_seconds();
    let latest_allowed = now.saturating_add(FUTURE_SKEW_SECONDS);
    let mut entries = HashMap::with_capacity(count);
    for _ in 0..count {
        let key = CacheKey::from_digest(decoder.array_32()?);
        let last_used_at = decoder.u64()?;
        let lru_order = decoder.u64()?;
        let value_length = usize::try_from(decoder.u32()?)
            .map_err(|_| invalid_data("cache value length is invalid"))?;
        if value_length > MAX_VALUE_BYTES {
            return Err(invalid_data("cache value exceeds size limit"));
        }
        if last_used_at > latest_allowed {
            return Err(invalid_data("cache timestamp is invalid"));
        }

        let value = decoder.take(value_length)?.to_vec();
        validate_value(&value)?;
        let entry = Entry {
            value: Arc::from(value),
            last_used_at,
            persisted_last_used_at: last_used_at,
            lru_order,
        };
        if entries.insert(key, entry).is_some() {
            return Err(invalid_data("cache file contains duplicate keys"));
        }
    }
    if !decoder.is_empty() {
        return Err(invalid_data("cache file contains trailing data"));
    }
    Ok(entries)
}

fn encode(output: &mut impl Write, entries: &HashMap<CacheKey, Entry>) -> io::Result<()> {
    output.write_all(&MAGIC)?;
    output.write_all(&CACHE_FORMAT_VERSION.to_be_bytes())?;
    output.write_all(
        &u32::try_from(entries.len())
            .unwrap_or(u32::MAX)
            .to_be_bytes(),
    )?;
    for (key, entry) in entries {
        output.write_all(&key.bytes())?;
        output.write_all(&entry.last_used_at.to_be_bytes())?;
        output.write_all(&entry.lru_order.to_be_bytes())?;
        output.write_all(
            &u32::try_from(entry.value.len())
                .map_err(|_| invalid_data("cache value length is invalid"))?
                .to_be_bytes(),
        )?;
        output.write_all(&entry.value)?;
    }
    Ok(())
}

fn ensure_private_directory(path: &Path) -> io::Result<()> {
    fs::create_dir_all(path)?;
    let metadata = fs::symlink_metadata(path)?;
    if !metadata.file_type().is_dir() {
        return Err(invalid_data("cache directory is not a directory"));
    }
    fs::set_permissions(path, fs::Permissions::from_mode(0o700))
}

fn validate_private_file(path: &Path) -> io::Result<()> {
    let Some(parent) = path.parent() else {
        return Err(invalid_data("cache path has no parent"));
    };
    let directory = fs::symlink_metadata(parent)?;
    if !directory.file_type().is_dir() || directory.mode() & 0o077 != 0 {
        return Err(invalid_data("cache directory permissions are unsafe"));
    }
    let metadata = fs::symlink_metadata(path)?;
    if !metadata.file_type().is_file() || metadata.mode() & 0o077 != 0 {
        return Err(invalid_data("cache file permissions are unsafe"));
    }
    Ok(())
}

fn temporary_path(path: &Path) -> PathBuf {
    let mut name = path.as_os_str().to_os_string();
    name.push(format!(".tmp-{}", std::process::id()));
    PathBuf::from(name)
}

fn open_temporary(path: &Path) -> io::Result<File> {
    match OpenOptions::new().write(true).create_new(true).open(path) {
        Ok(file) => Ok(file),
        Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {
            fs::remove_file(path)?;
            OpenOptions::new().write(true).create_new(true).open(path)
        }
        Err(error) => Err(error),
    }
}

fn sync_directory(path: &Path) -> io::Result<()> {
    File::open(path)?.sync_all()
}

struct Decoder<'a> {
    bytes: &'a [u8],
    position: usize,
}

impl<'a> Decoder<'a> {
    const fn new(bytes: &'a [u8]) -> Self {
        Self { bytes, position: 0 }
    }

    fn take(&mut self, length: usize) -> io::Result<&'a [u8]> {
        let end = self
            .position
            .checked_add(length)
            .ok_or_else(|| invalid_data("cache file length overflow"))?;
        let value = self
            .bytes
            .get(self.position..end)
            .ok_or_else(|| invalid_data("cache file is truncated"))?;
        self.position = end;
        Ok(value)
    }

    fn array_32(&mut self) -> io::Result<[u8; 32]> {
        self.take(32)?
            .try_into()
            .map_err(|_| invalid_data("cache file is truncated"))
    }

    fn u16(&mut self) -> io::Result<u16> {
        self.take(2)?
            .try_into()
            .map(u16::from_be_bytes)
            .map_err(|_| invalid_data("cache file is truncated"))
    }

    fn u32(&mut self) -> io::Result<u32> {
        self.take(4)?
            .try_into()
            .map(u32::from_be_bytes)
            .map_err(|_| invalid_data("cache file is truncated"))
    }

    fn u64(&mut self) -> io::Result<u64> {
        self.take(8)?
            .try_into()
            .map(u64::from_be_bytes)
            .map_err(|_| invalid_data("cache file is truncated"))
    }

    fn is_empty(&self) -> bool {
        self.position == self.bytes.len()
    }
}

fn invalid_data(message: &'static str) -> io::Error {
    io::Error::new(io::ErrorKind::InvalidData, message)
}

#[cfg(test)]
mod tests {
    use std::collections::HashMap;
    use std::fs;
    use std::os::unix::fs::PermissionsExt as _;
    use std::path::{Path, PathBuf};
    use std::sync::Arc;
    use std::sync::atomic::{AtomicU64, Ordering};

    use super::{
        CACHE_FORMAT_VERSION, CacheKey, Entry, FUTURE_SKEW_SECONDS, MAGIC, decode, encode, load,
        now_epoch_seconds, save,
    };

    static SEQUENCE: AtomicU64 = AtomicU64::new(0);

    struct TestDirectory(PathBuf);

    impl TestDirectory {
        fn new() -> Self {
            let sequence = SEQUENCE.fetch_add(1, Ordering::Relaxed);
            let path = std::env::temp_dir().join(format!(
                "ztheme-cache-disk-test-{}-{sequence}",
                std::process::id()
            ));
            fs::create_dir_all(&path).unwrap();
            fs::set_permissions(&path, fs::Permissions::from_mode(0o700)).unwrap();
            Self(path)
        }

        fn path(&self) -> &Path {
            &self.0
        }
    }

    impl Drop for TestDirectory {
        fn drop(&mut self) {
            let _ = fs::remove_dir_all(&self.0);
        }
    }

    fn entry(value: &[u8], last_used_at: u64, order: u64) -> Entry {
        Entry {
            value: Arc::from(value),
            last_used_at,
            persisted_last_used_at: last_used_at,
            lru_order: order,
        }
    }

    fn encoded(entries: &HashMap<CacheKey, Entry>) -> Vec<u8> {
        let mut output = Vec::new();
        encode(&mut output, entries).unwrap();
        output
    }

    #[test]
    fn save_load_preserves_entries_lru_and_private_permissions() {
        let directory = TestDirectory::new();
        let path = directory.path().join("nested/cache/runtime-v2.bin");
        let entries = HashMap::from([
            (CacheKey::from_value(1), entry(b"first", 100, 4)),
            (CacheKey::from_value(2), entry(b"second", 200, 9)),
        ]);

        save(&path, &entries).unwrap();
        let loaded = load(&path).unwrap();

        assert_eq!(loaded.len(), 2);
        for (key, expected) in &entries {
            let actual = &loaded[key];
            assert_eq!(actual.value.as_ref(), expected.value.as_ref());
            assert_eq!(actual.last_used_at, expected.last_used_at);
            assert_eq!(
                actual.persisted_last_used_at,
                expected.persisted_last_used_at
            );
            assert_eq!(actual.lru_order, expected.lru_order);
        }
        assert_eq!(
            fs::metadata(path.parent().unwrap())
                .unwrap()
                .permissions()
                .mode()
                & 0o777,
            0o700
        );
        assert_eq!(
            fs::metadata(&path).unwrap().permissions().mode() & 0o777,
            0o600
        );
    }

    #[test]
    fn old_entries_are_read_without_wall_clock_expiry() {
        let entries = HashMap::from([(CacheKey::from_value(1), entry(b"old", 1, 1))]);
        let loaded = decode(&encoded(&entries)).unwrap();
        assert_eq!(&*loaded[&CacheKey::from_value(1)].value, b"old");
    }

    #[test]
    fn malformed_cache_files_are_rejected() {
        let entries = HashMap::from([(CacheKey::from_value(1), entry(b"value", 1, 1))]);
        let valid = encoded(&entries);
        for length in 0..valid.len() {
            assert!(
                decode(&valid[..length]).is_err(),
                "accepted length {length}"
            );
        }
        let mut trailing = valid.clone();
        trailing.push(0);
        assert!(decode(&trailing).is_err());
        let mut bad_magic = valid.clone();
        bad_magic[0] ^= 1;
        assert!(decode(&bad_magic).is_err());
        let mut bad_version = valid.clone();
        bad_version[MAGIC.len()..MAGIC.len() + 2]
            .copy_from_slice(&CACHE_FORMAT_VERSION.saturating_add(1).to_be_bytes());
        assert!(decode(&bad_version).is_err());
    }

    #[test]
    fn duplicate_cache_keys_are_rejected() {
        let entries = HashMap::from([(CacheKey::from_value(1), entry(b"value", 1, 1))]);
        let valid = encoded(&entries);
        let mut duplicate = valid.clone();
        duplicate[6..10].copy_from_slice(&2_u32.to_be_bytes());
        duplicate.extend_from_slice(&valid[10..]);
        assert!(decode(&duplicate).is_err());
    }

    #[test]
    fn future_timestamps_are_rejected() {
        let entries = HashMap::from([(CacheKey::from_value(1), entry(b"value", 1, 1))]);
        let mut future = encoded(&entries);
        let timestamp = now_epoch_seconds().saturating_add(FUTURE_SKEW_SECONDS + 1);
        future[42..50].copy_from_slice(&timestamp.to_be_bytes());
        assert!(decode(&future).is_err());
    }
}