use std::fs;
use std::path::{Path, PathBuf};
use anyhow::{anyhow, Context as _};
use super::release::binary_name;
const MAX_EXTRACTED_BYTES: u64 = 256 * 1024 * 1024;
fn write_member<R: std::io::Read>(
dest: &Path,
declared: u64,
reader: R,
limit: u64,
) -> anyhow::Result<()> {
if declared > limit {
return Err(anyhow!(
"the archive's {} is {declared} bytes, which is too large to be the pmpx binary",
binary_name(std::env::consts::OS)
));
}
let mut out =
fs::File::create(dest).with_context(|| format!("cannot write {}", dest.display()))?;
let mut limited = reader.take(limit + 1);
let copied = std::io::copy(&mut limited, &mut out).context("cannot read the archive entry")?;
if copied > limit {
drop(out);
let _ = fs::remove_file(dest);
return Err(anyhow!(
"the entry for {} is larger than the {limit} byte limit",
dest.display()
));
}
out.sync_all()
.with_context(|| format!("cannot flush {}", dest.display()))?;
Ok(())
}
pub(super) fn extract(is_zip: bool, archive: &[u8], dir: &Path) -> anyhow::Result<PathBuf> {
let member = binary_name(std::env::consts::OS);
let extracted = if is_zip {
extract_from_zip(archive, member, dir)?
} else {
extract_from_tar_gz(archive, member, dir)?
};
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt as _;
fs::set_permissions(&extracted, fs::Permissions::from_mode(0o755))
.context("cannot make the new binary executable")?;
}
Ok(extracted)
}
pub(super) fn extract_from_tar_gz(
archive: &[u8],
member: &str,
dir: &Path,
) -> anyhow::Result<PathBuf> {
let decoder = flate2::read::GzDecoder::new(archive);
let mut tar = tar::Archive::new(decoder);
for entry in tar
.entries()
.context("the release archive is not a tar file")?
{
let mut entry = entry.context("the release archive is damaged")?;
let path = entry
.path()
.context("an entry in the archive has no name")?;
if path.file_name().and_then(|name| name.to_str()) != Some(member) {
continue;
}
if !entry.header().entry_type().is_file() {
continue;
}
let dest = dir.join(member);
write_member(
&dest,
entry.header().size().unwrap_or(0),
&mut entry,
MAX_EXTRACTED_BYTES,
)?;
return Ok(dest);
}
Err(anyhow!("the release archive does not contain {member}"))
}
pub(super) fn extract_from_zip(
archive: &[u8],
member: &str,
dir: &Path,
) -> anyhow::Result<PathBuf> {
let reader = std::io::Cursor::new(archive);
let mut zip = zip::ZipArchive::new(reader).context("the release archive is not a zip file")?;
for index in 0..zip.len() {
let mut file = zip
.by_index(index)
.context("the release archive is damaged")?;
if !file.is_file() {
continue;
}
let name = file.name().to_string();
if Path::new(&name).file_name().and_then(|name| name.to_str()) != Some(member) {
continue;
}
let dest = dir.join(member);
write_member(&dest, file.size(), &mut file, MAX_EXTRACTED_BYTES)?;
return Ok(dest);
}
Err(anyhow!("the release archive does not contain {member}"))
}
#[cfg(test)]
mod tests {
use super::*;
fn tar_gz_with(entries: &[(&str, &[u8])]) -> Vec<u8> {
let mut builder = tar::Builder::new(Vec::new());
for (name, body) in entries {
let mut header = tar::Header::new_gnu();
header.set_size(body.len() as u64);
header.set_mode(0o755);
header.set_cksum();
builder.append_data(&mut header, name, *body).unwrap();
}
let tar = builder.into_inner().unwrap();
let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
std::io::Write::write_all(&mut encoder, &tar).unwrap();
encoder.finish().unwrap()
}
fn zip_with(entries: &[(&str, &[u8])]) -> Vec<u8> {
use std::io::Write as _;
let mut writer = zip::ZipWriter::new(std::io::Cursor::new(Vec::new()));
let options: zip::write::FileOptions<'_, ()> =
zip::write::FileOptions::default().compression_method(zip::CompressionMethod::Deflated);
for (name, body) in entries {
writer.start_file(*name, options).unwrap();
writer.write_all(body).unwrap();
}
writer.finish().unwrap().into_inner()
}
#[test]
fn the_binary_comes_out_of_a_tarball() {
let archive = tar_gz_with(&[("LICENSE", b"licence"), ("pmpx", b"the new pmpx")]);
let tmp = tempfile::tempdir().unwrap();
let got = extract_from_tar_gz(&archive, "pmpx", tmp.path()).expect("should extract");
assert_eq!(std::fs::read(&got).unwrap(), b"the new pmpx");
assert_eq!(got.file_name().unwrap(), "pmpx");
}
#[test]
fn the_binary_comes_out_of_a_zip() {
let archive = zip_with(&[("LICENSE", b"licence"), ("pmpx.exe", b"the new pmpx.exe")]);
let tmp = tempfile::tempdir().unwrap();
let got = extract_from_zip(&archive, "pmpx.exe", tmp.path()).expect("should extract");
assert_eq!(std::fs::read(&got).unwrap(), b"the new pmpx.exe");
assert_eq!(got.file_name().unwrap(), "pmpx.exe");
}
#[test]
fn an_archive_without_the_binary_is_refused() {
let archive = tar_gz_with(&[("LICENSE", b"licence")]);
let tmp = tempfile::tempdir().unwrap();
let error = extract_from_tar_gz(&archive, "pmpx", tmp.path()).unwrap_err();
assert!(
error.to_string().contains("does not contain pmpx"),
"{error}"
);
}
#[test]
fn something_that_is_not_an_archive_is_refused() {
let tmp = tempfile::tempdir().unwrap();
assert!(extract_from_tar_gz(b"not a tarball", "pmpx", tmp.path()).is_err());
assert!(extract_from_zip(b"not a zip", "pmpx.exe", tmp.path()).is_err());
}
#[test]
fn the_destination_does_not_follow_the_archive_entry_name() {
let archive = tar_gz_with(&[("nested/dir/pmpx", b"somewhere else")]);
let tmp = tempfile::tempdir().unwrap();
let got = extract_from_tar_gz(&archive, "pmpx", tmp.path()).expect("should extract");
assert_eq!(got.parent().unwrap(), tmp.path());
assert_eq!(std::fs::read(&got).unwrap(), b"somewhere else");
}
fn tar_gz_with_a_link_named(member: &str) -> Vec<u8> {
let mut builder = tar::Builder::new(Vec::new());
let mut header = tar::Header::new_gnu();
header.set_entry_type(tar::EntryType::Symlink);
header.set_size(0);
header.set_mode(0o777);
header.set_link_name("/etc/passwd").unwrap();
header.set_cksum();
builder
.append_data(&mut header, member, std::io::empty())
.unwrap();
let tar = builder.into_inner().unwrap();
let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
std::io::Write::write_all(&mut encoder, &tar).unwrap();
encoder.finish().unwrap()
}
#[test]
fn a_link_entry_named_like_the_binary_is_refused() {
let archive = tar_gz_with_a_link_named("pmpx");
let tmp = tempfile::tempdir().unwrap();
let error = extract_from_tar_gz(&archive, "pmpx", tmp.path()).unwrap_err();
assert!(
error.to_string().contains("does not contain pmpx"),
"{error}"
);
assert!(
!tmp.path().join("pmpx").exists(),
"a link entry must not create anything at the destination"
);
}
#[test]
fn a_directory_entry_named_like_the_binary_is_refused() {
let mut writer = zip::ZipWriter::new(std::io::Cursor::new(Vec::new()));
let options: zip::write::FileOptions<'_, ()> =
zip::write::FileOptions::default().compression_method(zip::CompressionMethod::Stored);
writer.add_directory("pmpx.exe", options).unwrap();
let archive = writer.finish().unwrap().into_inner();
let tmp = tempfile::tempdir().unwrap();
let error = extract_from_zip(&archive, "pmpx.exe", tmp.path()).unwrap_err();
assert!(
error.to_string().contains("does not contain pmpx.exe"),
"{error}"
);
assert!(!tmp.path().join("pmpx.exe").is_file());
}
#[test]
fn a_member_that_claims_to_be_enormous_is_refused() {
let tmp = tempfile::tempdir().unwrap();
let dest = tmp.path().join("pmpx");
let error = write_member(&dest, 100, &b"short"[..], 16).unwrap_err();
assert!(error.to_string().contains("too large"), "{error}");
assert!(!dest.exists(), "nothing should have been created");
}
#[test]
fn a_member_larger_than_it_claims_is_refused() {
let tmp = tempfile::tempdir().unwrap();
let dest = tmp.path().join("pmpx");
let ten_bytes = b"0123456789";
let error = write_member(&dest, 1, &ten_bytes[..], 4).unwrap_err();
assert!(error.to_string().contains("larger than"), "{error}");
assert!(!dest.exists(), "the partial file must not be left behind");
}
#[test]
fn a_member_within_the_limit_is_written() {
let tmp = tempfile::tempdir().unwrap();
let dest = tmp.path().join("pmpx");
write_member(&dest, 5, &b"hello"[..], 16).unwrap();
assert_eq!(std::fs::read(&dest).unwrap(), b"hello");
}
}