use std::collections::HashMap;
use std::io;
use std::path::{Component, Path, PathBuf};
use crate::EngineError;
pub fn write_tar_sync(root_dir: &Path, out: &mut impl io::Write) -> Result<(), io::Error> {
let mut builder = tar::Builder::new(out);
builder.follow_symlinks(false);
append_dir_dedup(&mut builder, root_dir)?;
builder.finish()?;
Ok(())
}
fn append_dir_dedup<W: io::Write>(
builder: &mut tar::Builder<W>,
root_dir: &Path,
) -> Result<(), io::Error> {
let mut seen: HashMap<(u64, u64), PathBuf> = HashMap::new();
for entry in walkdir::WalkDir::new(root_dir)
.follow_links(false)
.sort_by_file_name()
.into_iter()
.filter_map(std::result::Result::ok)
{
let path = entry.path();
let relative = path.strip_prefix(root_dir).unwrap_or(path);
if entry.file_type().is_dir() {
let is_root = relative.as_os_str().is_empty() || relative == Path::new(".");
if !is_root {
builder.append_dir(relative, path)?;
}
} else if entry.file_type().is_file() {
let mut is_dup = false;
#[cfg_attr(not(unix), allow(unused_variables))]
if let Ok(meta) = entry.metadata() {
#[cfg(unix)]
let nlink = {
use std::os::unix::fs::MetadataExt;
meta.nlink()
};
#[cfg(not(unix))]
let nlink: u64 = 1;
if nlink > 1 {
#[cfg(unix)]
let key = {
use std::os::unix::fs::MetadataExt;
(meta.dev(), meta.ino())
};
#[cfg(not(unix))]
let key = (0u64, 0u64);
if let Some(first) = seen.get(&key) {
let link_target = relative_path_from_to(
relative.parent().unwrap_or(Path::new(".")),
first,
);
let mut header = tar::Header::new_gnu();
header.set_entry_type(tar::EntryType::Link);
header.set_size(0);
builder.append_link(&mut header, relative, link_target)?;
is_dup = true;
} else {
seen.insert(key, relative.to_path_buf());
}
}
}
if !is_dup {
builder.append_path_with_name(path, relative)?;
}
}
}
Ok(())
}
fn relative_path_from_to(from_dir: &Path, to_path: &Path) -> PathBuf {
let from: Vec<_> = from_dir.components().collect();
let to: Vec<_> = to_path.components().collect();
let common = from.iter().zip(to.iter()).take_while(|(a, b)| a == b).count();
let mut result = PathBuf::new();
for _ in common..from.len() {
result.push("..");
}
for comp in &to[common..] {
result.push(comp);
}
result
}
pub fn extract_tar_sync(input: impl io::Read, output_dir: &Path) -> Result<(), EngineError> {
std::fs::create_dir_all(output_dir).map_err(EngineError::Io)?;
let output_dir = std::fs::canonicalize(output_dir).map_err(EngineError::Io)?;
let mut archive = tar::Archive::new(input);
let mut pending_hard_links = Vec::new();
for entry in archive.entries().map_err(EngineError::Io)? {
let mut entry = entry.map_err(EngineError::Io)?;
let entry_path = entry.path().map_err(EngineError::Io)?;
let entry_type = entry.header().entry_type();
if entry_type.is_symlink() {
return Err(EngineError::PathTraversal);
}
if entry_path.is_absolute() {
return Err(EngineError::PathTraversal);
}
for component in entry_path.components() {
if matches!(component, Component::ParentDir) {
return Err(EngineError::PathTraversal);
}
}
let dest = output_dir.join(&entry_path);
if !dest.starts_with(&output_dir) {
return Err(EngineError::PathTraversal);
}
if entry_type.is_hard_link() {
let link_target =
entry.link_name().map_err(EngineError::Io)?.ok_or(EngineError::PathTraversal)?;
if link_target.is_absolute() {
return Err(EngineError::PathTraversal);
}
let link_dest = if let Some(parent) = dest.parent() {
parent.join(&link_target)
} else {
output_dir.join(&link_target)
};
let normalized = normalize_path(&link_dest);
if !normalized.starts_with(&output_dir) {
return Err(EngineError::PathTraversal);
}
pending_hard_links.push((dest, normalized));
continue;
}
if let Some(parent) = dest.parent() {
std::fs::create_dir_all(parent).map_err(EngineError::Io)?;
}
entry.unpack(&dest).map_err(EngineError::Io)?;
}
for (dest, normalized) in pending_hard_links {
if !normalized.is_file() {
return Err(EngineError::PathTraversal);
}
if let Some(parent) = dest.parent() {
std::fs::create_dir_all(parent).map_err(EngineError::Io)?;
}
std::fs::hard_link(&normalized, &dest).map_err(EngineError::Io)?;
}
Ok(())
}
fn normalize_path(path: &Path) -> PathBuf {
let mut result = PathBuf::new();
let mut depth: i32 = 0;
for component in path.components() {
match component {
Component::ParentDir => {
if depth > 0 {
result.pop();
depth -= 1;
} else {
result.push("..");
}
},
Component::CurDir => {},
c => {
result.push(c);
depth += 1;
},
}
}
result
}
pub fn estimate_dir_size(root_dir: &Path) -> u64 {
#[cfg_attr(not(unix), allow(unused_mut, unused_variables))]
let mut seen: std::collections::HashSet<(u64, u64)> = std::collections::HashSet::new();
walkdir::WalkDir::new(root_dir)
.follow_links(false)
.into_iter()
.filter_map(std::result::Result::ok)
.filter(|e| e.file_type().is_file())
.filter_map(|e| {
let meta = e.metadata().ok()?;
#[cfg(unix)]
{
use std::os::unix::fs::MetadataExt;
if meta.nlink() > 1 && !seen.insert((meta.dev(), meta.ino())) {
return None; }
}
Some(meta.len())
})
.sum()
}
#[cfg(test)]
mod tests {
use std::fs;
use std::io::Cursor;
use std::time::{SystemTime, UNIX_EPOCH};
use super::*;
fn temp_output(name: &str) -> std::path::PathBuf {
let unique = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("system clock is before UNIX_EPOCH")
.as_nanos();
std::env::temp_dir().join(format!("hayate-{name}-{}-{unique}", std::process::id()))
}
#[test]
fn extract_tar_creates_output_root() {
let mut archive = Vec::new();
{
let mut builder = tar::Builder::new(&mut archive);
let bytes = b"hello";
let mut header = tar::Header::new_gnu();
header.set_path("nested/file.txt").unwrap();
header.set_size(bytes.len() as u64);
header.set_cksum();
builder.append(&header, bytes.as_slice()).unwrap();
builder.finish().unwrap();
}
let out = temp_output("extract-root");
let result = extract_tar_sync(Cursor::new(archive), &out);
assert!(result.is_ok());
assert_eq!(fs::read_to_string(out.join("nested/file.txt")).unwrap(), "hello");
fs::remove_dir_all(out).unwrap();
}
#[test]
fn extract_tar_rejects_symlink_entries() {
let mut archive = Vec::new();
{
let mut builder = tar::Builder::new(&mut archive);
let mut header = tar::Header::new_gnu();
header.set_entry_type(tar::EntryType::Symlink);
header.set_path("link").unwrap();
header.set_link_name("../outside").unwrap();
header.set_size(0);
header.set_cksum();
builder.append(&header, Cursor::new(Vec::new())).unwrap();
builder.finish().unwrap();
}
let out = temp_output("reject-link");
let result = extract_tar_sync(Cursor::new(archive), &out);
assert!(matches!(result, Err(EngineError::PathTraversal)));
let _ = fs::remove_dir_all(out);
}
#[test]
fn hard_link_roundtrip() {
let src = temp_output("hardlink-src");
let dst = temp_output("hardlink-dst");
let sub = src.join("sub");
fs::create_dir_all(&sub).unwrap();
fs::write(sub.join("original.txt"), b"same content").unwrap();
fs::hard_link(sub.join("original.txt"), sub.join("link.txt")).unwrap();
let mut archive = Vec::new();
write_tar_sync(&src, &mut archive).unwrap();
extract_tar_sync(Cursor::new(&archive), &dst).unwrap();
let orig = dst.join("sub/original.txt");
let link = dst.join("sub/link.txt");
assert!(orig.is_file());
assert!(link.is_file());
assert_eq!(fs::read_to_string(&orig).unwrap(), "same content");
assert_eq!(fs::read_to_string(&link).unwrap(), "same content");
#[cfg(unix)]
{
use std::os::unix::fs::MetadataExt;
let orig_meta = fs::metadata(&orig).unwrap();
let link_meta = fs::metadata(&link).unwrap();
assert_eq!(orig_meta.ino(), link_meta.ino());
assert!(orig_meta.nlink() >= 2);
}
let _ = fs::remove_dir_all(src);
let _ = fs::remove_dir_all(dst);
}
#[test]
fn hard_link_rejects_path_traversal() {
let out = temp_output("hardlink-reject");
let mut archive = Vec::new();
{
let mut builder = tar::Builder::new(&mut archive);
let mut header = tar::Header::new_gnu();
header.set_path("innocent.txt").unwrap();
header.set_size(5);
header.set_cksum();
builder.append(&header, Cursor::new(b"hello")).unwrap();
let mut link_header = tar::Header::new_gnu();
link_header.set_entry_type(tar::EntryType::Link);
link_header.set_path("escape").unwrap();
link_header.set_link_name("../../../etc/passwd").unwrap();
link_header.set_size(0);
link_header.set_cksum();
builder.append(&link_header, Cursor::new(Vec::new())).unwrap();
builder.finish().unwrap();
}
let result = extract_tar_sync(Cursor::new(&archive), &out);
assert!(matches!(result, Err(EngineError::PathTraversal)));
let _ = fs::remove_dir_all(out);
}
#[test]
fn estimate_dir_size_dedup_hard_links() {
let dir = temp_output("estimate-dedup");
fs::create_dir_all(&dir).unwrap();
fs::write(dir.join("a.txt"), b"hello world").unwrap(); fs::hard_link(dir.join("a.txt"), dir.join("b.txt")).unwrap();
fs::hard_link(dir.join("a.txt"), dir.join("c.txt")).unwrap();
let total = estimate_dir_size(&dir);
#[cfg(unix)]
assert_eq!(total, 11);
#[cfg(not(unix))]
assert_eq!(total, 33);
fs::remove_dir_all(dir).unwrap();
}
#[test]
fn relative_path_from_to_same_dir() {
let result = relative_path_from_to(
Path::new("target/debug/incremental"),
Path::new("target/debug/incremental/hayate-abc"),
);
assert_eq!(result, PathBuf::from("hayate-abc"));
}
#[test]
fn relative_path_from_to_parent_dir() {
let result =
relative_path_from_to(Path::new("target/debug/x"), Path::new("target/debug/y/file"));
assert_eq!(result, PathBuf::from("../y/file"));
}
#[test]
fn relative_path_from_to_nested() {
let result = relative_path_from_to(Path::new("a/b/c"), Path::new("a/d/e"));
assert_eq!(result, PathBuf::from("../../d/e"));
}
#[test]
fn normalize_path_removes_dot() {
assert_eq!(normalize_path(Path::new("foo/./bar")), PathBuf::from("foo/bar"));
}
#[test]
fn normalize_path_resolves_dotdot() {
assert_eq!(normalize_path(Path::new("foo/bar/../baz")), PathBuf::from("foo/baz"));
}
#[test]
fn normalize_path_does_not_escape_root() {
assert_eq!(normalize_path(Path::new("foo/../../baz")), PathBuf::from("../baz"));
}
}