use std::io::{Read, Write};
use crate::artifact::SupportedFormat;
const MAX_UNPACKED_BYTES: u64 = 2 * 1024 * 1024 * 1024;
#[derive(Debug, thiserror::Error)]
pub enum ArchiveError {
#[error("failed to read the archive fetched from '{uri}': {reason}")]
Malformed { uri: String, reason: String },
#[error("the archive fetched from '{uri}' contains more than one file")]
NonZeroFiles { uri: String },
#[error("the archive fetched from '{uri}' contains no files at all")]
Empty { uri: String },
#[error("the archive fetched from '{uri}' unpacks to more than {limit} bytes")]
TooLarge { uri: String, limit: u64 },
#[error("failed to write the file unpacked from '{uri}': {reason}")]
Unwritable { uri: String, reason: String },
}
pub fn extract(
body: &[u8],
format: SupportedFormat,
uri: &str,
out: &mut impl Write,
) -> Result<(), ArchiveError> {
extract_bounded(body, format, uri, out, MAX_UNPACKED_BYTES)
}
fn extract_bounded(
body: &[u8],
format: SupportedFormat,
uri: &str,
out: &mut impl Write,
limit: u64,
) -> Result<(), ArchiveError> {
let mut stream = match format {
SupportedFormat::TarGz => flate2::read::MultiGzDecoder::new(body).take(limit + 1),
};
let result = from_tar(&mut stream, uri, out);
if stream.limit() == 0 {
return Err(ArchiveError::TooLarge { uri: uri.to_string(), limit });
}
result
}
fn from_tar<R: Read>(stream: R, uri: &str, out: &mut impl Write) -> Result<(), ArchiveError> {
let malformed = |err: std::io::Error| ArchiveError::Malformed {
uri: uri.to_string(),
reason: err.to_string(),
};
let mut tar = tar::Archive::new(stream);
tar.set_ignore_zeros(true);
let entries = tar.entries().map_err(malformed)?;
let mut found = false;
let mut buffer = [0u8; 64 * 1024];
for entry in entries {
let mut entry = entry.map_err(malformed)?;
if !entry.header().entry_type().is_file() {
continue;
}
if found {
return Err(ArchiveError::NonZeroFiles { uri: uri.to_string() });
}
found = true;
loop {
let read = entry.read(&mut buffer).map_err(malformed)?;
if read == 0 {
break;
}
out.write_all(&buffer[..read]).map_err(|err| ArchiveError::Unwritable {
uri: uri.to_string(),
reason: err.to_string(),
})?;
}
}
let mut rest = tar.into_inner();
std::io::copy(&mut rest, &mut std::io::sink()).map_err(malformed)?;
if !found {
return Err(ArchiveError::Empty { uri: uri.to_string() });
}
Ok(())
}
#[cfg(test)]
mod tests {
use flate2::{Compression, write::GzEncoder};
use super::*;
const URI: &str = "https://example.invalid/artifact.tar.gz";
const TRAILER: usize = 8;
fn extracted(body: &[u8]) -> Result<Vec<u8>, ArchiveError> {
let mut out = Vec::new();
extract(body, SupportedFormat::TarGz, URI, &mut out).map(|()| out)
}
fn extracted_under(limit: u64, body: &[u8]) -> Result<Vec<u8>, ArchiveError> {
let mut out = Vec::new();
extract_bounded(body, SupportedFormat::TarGz, URI, &mut out, limit).map(|()| out)
}
fn tar(entries: &[(&str, &[u8])]) -> Vec<u8> {
let mut tar = tar::Builder::new(Vec::new());
for (path, contents) in entries {
let mut header = tar::Header::new_gnu();
header.set_size(contents.len() as u64);
header.set_mode(0o644);
header.set_cksum();
tar.append_data(&mut header, path, *contents).unwrap();
}
tar.into_inner().unwrap()
}
fn tarball(entries: &[(&str, &[u8])]) -> Vec<u8> {
gzipped(tar(entries))
}
fn gzipped(plain: Vec<u8>) -> Vec<u8> {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(&plain).unwrap();
encoder.finish().unwrap()
}
fn directory(path: &str) -> tar::Header {
let mut header = tar::Header::new_gnu();
header.set_entry_type(tar::EntryType::Directory);
header.set_size(0);
header.set_mode(0o755);
header.set_path(path).unwrap();
header.set_cksum();
header
}
fn nested() -> Vec<u8> {
let mut tar = tar::Builder::new(Vec::new());
tar.append(&directory("miden-vm-aarch64-apple-darwin/"), std::io::empty())
.unwrap();
let mut header = tar::Header::new_gnu();
header.set_size(6);
header.set_mode(0o755);
header.set_cksum();
tar.append_data(&mut header, "miden-vm-aarch64-apple-darwin/miden-vm", &b"binary"[..])
.unwrap();
gzipped(tar.into_inner().unwrap())
}
#[test]
fn the_sole_file_is_extracted() {
let extracted = extracted(&nested()).expect("should extract");
assert_eq!(extracted, b"binary", "the directory entry must not count as a file");
}
#[test]
fn a_file_larger_than_the_copy_buffer_arrives_whole() {
let contents: Vec<u8> = (0..300_000u32).map(|byte| byte as u8).collect();
let extracted = extracted(&tarball(&[("big", &contents)])).expect("should extract");
assert_eq!(extracted, contents, "every read must reach the writer, in order");
}
#[test]
fn several_files_is_an_error() {
let err = extracted(&tarball(&[("a", b"one"), ("b", b"two")])).expect_err("must fail");
assert!(matches!(err, ArchiveError::NonZeroFiles { .. }), "{err}");
assert!(err.to_string().contains(URI), "the message must name the source: {err}");
}
#[test]
fn files_after_a_tar_end_marker_are_still_counted() {
let mut joined = tar(&[("a", b"one")]);
joined.extend(tar(&[("b", b"two")]));
let err = extracted(&gzipped(joined)).expect_err("must fail");
assert!(matches!(err, ArchiveError::NonZeroFiles { .. }), "{err}");
}
#[test]
fn an_archive_with_no_files_is_an_error() {
let mut tar = tar::Builder::new(Vec::new());
tar.append(&directory("empty/"), std::io::empty()).unwrap();
let err = extracted(&gzipped(tar.into_inner().unwrap())).expect_err("must fail");
assert!(matches!(err, ArchiveError::Empty { .. }), "{err}");
}
#[test]
fn a_missing_gzip_trailer_is_an_error() {
let mut body = nested();
body.truncate(body.len() - TRAILER);
let err = extracted(&body).expect_err("must fail");
assert!(matches!(err, ArchiveError::Malformed { .. }), "{err}");
assert!(err.to_string().contains(URI), "the message must name the source: {err}");
}
#[test]
fn a_gzip_checksum_that_does_not_match_is_an_error() {
let mut body = nested();
let crc = body.len() - TRAILER;
body[crc] ^= 0xff;
let err = extracted(&body).expect_err("must fail");
assert!(matches!(err, ArchiveError::Malformed { .. }), "{err}");
}
#[test]
fn a_later_gzip_member_checksum_must_match() {
let mut body = nested();
body.extend(gzipped(vec![0; 1024]));
assert_eq!(extracted(&body).expect("every valid member should be accepted"), b"binary");
let crc = body.len() - TRAILER;
body[crc] ^= 0xff;
let err = extracted(&body).expect_err("must fail");
assert!(matches!(err, ArchiveError::Malformed { .. }), "{err}");
}
#[test]
fn garbage_is_reported_as_a_malformed_archive() {
let err = extracted(b"<html>not a tarball</html>").expect_err("must fail");
assert!(matches!(err, ArchiveError::Malformed { .. }), "{err}");
assert!(err.to_string().contains(URI), "the message must name the source: {err}");
}
#[test]
fn a_writer_that_fails_is_reported_as_such() {
struct Full;
impl Write for Full {
fn write(&mut self, _: &[u8]) -> std::io::Result<usize> {
Err(std::io::Error::other("no space left on device"))
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
let err =
extract(&nested(), SupportedFormat::TarGz, URI, &mut Full).expect_err("must fail");
assert!(matches!(err, ArchiveError::Unwritable { .. }), "{err}");
assert!(err.to_string().contains(URI), "the message must name the source: {err}");
}
#[test]
fn an_archive_that_unpacks_past_the_limit_is_refused() {
let body = tarball(&[("big", &vec![0u8; 512 * 1024])]);
assert!(body.len() < 4096, "the compressed size must say nothing about the unpacked one");
let limit = 64 * 1024;
let err = extracted_under(limit, &body).expect_err("must fail");
assert!(matches!(err, ArchiveError::TooLarge { .. }), "{err}");
assert!(err.to_string().contains(URI), "the message must name the source: {err}");
assert!(
err.to_string().contains(&limit.to_string()),
"the message must name the limit: {err}"
);
}
#[test]
fn an_archive_that_ends_on_the_limit_is_accepted() {
let body = tarball(&[("snug", b"contents")]);
let mut plain = Vec::new();
flate2::read::MultiGzDecoder::new(&body[..]).read_to_end(&mut plain).unwrap();
let extracted = extracted_under(plain.len() as u64, &body)
.expect("a stream exactly the limit's length must be read");
assert_eq!(extracted, b"contents");
}
}