pub mod sanitize;
use crate::error::Error;
use sanitize::{BoundedReader, DEFAULT_MAX_TOTAL_UNCOMPRESSED_SIZE, DEFAULT_MAX_UNCOMPRESSED_SIZE};
use std::io::{Read, Seek};
#[derive(Debug)]
pub struct ZipContainer<R> {
archive: zip::ZipArchive<R>,
max_entry_size: u64,
max_total_size: u64,
total_read: u64,
}
impl<R: Read + Seek> ZipContainer<R> {
pub fn open_reader(reader: R) -> Result<Self, Error> {
let archive =
zip::ZipArchive::new(reader).map_err(|e| Error::InvalidPackage(e.to_string()))?;
for name in archive.file_names() {
sanitize::validate_entry_path(name)?;
}
Ok(Self {
archive,
max_entry_size: DEFAULT_MAX_UNCOMPRESSED_SIZE,
max_total_size: DEFAULT_MAX_TOTAL_UNCOMPRESSED_SIZE,
total_read: 0,
})
}
pub fn get_entry(
&mut self,
name: &str,
) -> Result<Option<BoundedReader<'_, impl Read + '_>>, Error> {
sanitize::validate_entry_path(name)?;
let Self {
archive,
total_read,
max_entry_size,
max_total_size,
..
} = self;
match archive.by_name(name) {
Ok(file) => Ok(Some(BoundedReader::new(
file,
*max_entry_size,
total_read,
*max_total_size,
))),
Err(zip::result::ZipError::FileNotFound) => Ok(None),
Err(e) => Err(Error::InvalidPackage(e.to_string())),
}
}
#[allow(dead_code)]
pub fn entry_names(&self) -> impl Iterator<Item = &str> {
self.archive.file_names()
}
}
impl<R> ZipContainer<R> {
pub(crate) fn with_max_entry_size(mut self, limit: u64) -> Self {
self.max_entry_size = limit;
self
}
pub(crate) fn with_max_total_size(mut self, limit: u64) -> Self {
self.max_total_size = limit;
self
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::{Cursor, Write};
fn build_zip(entries: &[(&str, &[u8])]) -> Vec<u8> {
let mut buf = Vec::new();
{
let mut writer = zip::ZipWriter::new(Cursor::new(&mut buf));
let options = zip::write::SimpleFileOptions::default()
.compression_method(zip::CompressionMethod::Deflated);
for (name, data) in entries {
writer.start_file(*name, options).unwrap();
writer.write_all(data).unwrap();
}
writer.finish().unwrap();
}
buf
}
#[test]
fn open_reader_succeeds_for_valid_xlsx_shaped_zip() {
let bytes = build_zip(&[
("[Content_Types].xml", b"<Types/>"),
("xl/workbook.xml", b"<workbook/>"),
]);
let container = ZipContainer::open_reader(Cursor::new(bytes)).unwrap();
let mut names: Vec<&str> = container.entry_names().collect();
names.sort_unstable();
assert_eq!(names, vec!["[Content_Types].xml", "xl/workbook.xml"]);
}
#[test]
fn open_reader_rejects_corrupt_zip_bytes() {
let err = ZipContainer::open_reader(Cursor::new(b"not a zip file".to_vec())).unwrap_err();
assert!(matches!(err, Error::InvalidPackage(_)));
}
#[test]
fn open_reader_rejects_archive_with_invalid_entry_name() {
let bytes = build_zip(&[("../evil", b"payload")]);
let err = ZipContainer::open_reader(Cursor::new(bytes)).unwrap_err();
assert!(matches!(err, Error::ZipSlipDetected { .. }));
}
#[test]
fn get_entry_returns_content_for_existing_name() {
let bytes = build_zip(&[("xl/workbook.xml", b"<workbook/>")]);
let mut container = ZipContainer::open_reader(Cursor::new(bytes)).unwrap();
let mut out = Vec::new();
container
.get_entry("xl/workbook.xml")
.unwrap()
.unwrap()
.read_to_end(&mut out)
.unwrap();
assert_eq!(out, b"<workbook/>");
}
#[test]
fn get_entry_returns_none_for_missing_name() {
let bytes = build_zip(&[("xl/workbook.xml", b"<workbook/>")]);
let mut container = ZipContainer::open_reader(Cursor::new(bytes)).unwrap();
assert!(container.get_entry("xl/missing.xml").unwrap().is_none());
}
#[test]
fn get_entry_rejects_malformed_name_even_if_absent() {
let bytes = build_zip(&[("xl/workbook.xml", b"<workbook/>")]);
let mut container = ZipContainer::open_reader(Cursor::new(bytes)).unwrap();
let result = container.get_entry("../etc/passwd");
assert!(matches!(result, Err(Error::ZipSlipDetected { .. })));
}
#[test]
fn get_entry_enforces_per_entry_size_limit() {
let bytes = build_zip(&[("xl/workbook.xml", &[0u8; 100])]);
let mut container = ZipContainer::open_reader(Cursor::new(bytes)).unwrap();
container.max_entry_size = 10;
let mut reader = container.get_entry("xl/workbook.xml").unwrap().unwrap();
let mut out = Vec::new();
assert!(reader.read_to_end(&mut out).is_err());
}
#[test]
fn total_read_accumulates_and_enforces_cumulative_limit() {
let bytes = build_zip(&[("xl/a.xml", &[0u8; 10]), ("xl/b.xml", &[0u8; 10])]);
let mut container = ZipContainer::open_reader(Cursor::new(bytes)).unwrap();
container.max_total_size = 15;
let mut out = Vec::new();
container
.get_entry("xl/a.xml")
.unwrap()
.unwrap()
.read_to_end(&mut out)
.unwrap();
assert_eq!(container.total_read, 10);
let mut reader = container.get_entry("xl/b.xml").unwrap().unwrap();
let mut out2 = Vec::new();
assert!(reader.read_to_end(&mut out2).is_err());
}
}