Skip to main content

conspire/io/zip/
mod.rs

1#[cfg(test)]
2mod test;
3
4use crate::io::Write;
5use std::{
6    fs::File,
7    io::{self, BufReader, BufWriter, Error, ErrorKind, Read, Result},
8    path::Path,
9    str::from_utf8,
10};
11
12/// One member of a ZIP archive.
13///
14/// `name` is the path within the archive; `data` is the member's uncompressed
15/// contents.
16pub struct ZipEntry {
17    pub name: String,
18    pub data: Vec<u8>,
19}
20
21/// An in-memory ZIP archive.
22///
23/// Holds an ordered list of [`ZipEntry`] members. Read with [`Zip::read`] (look
24/// entries up by name with [`Zip::entry`]) and written through the
25/// [`Write`](crate::io::Write) trait. Only stored (uncompressed) entries are
26/// supported, which is what NumPy `.npz` files use.
27pub struct Zip {
28    pub entries: Vec<ZipEntry>,
29}
30
31impl<P: AsRef<Path>> Write<P> for Zip {
32    type Error = Error;
33    fn write(&self, path: P) -> Result<()> {
34        self.write_to(&mut BufWriter::new(File::create(path)?))
35    }
36}
37
38impl Zip {
39    fn write_to<W: io::Write>(&self, file: &mut W) -> Result<()> {
40        let mut offset: u32 = 0;
41        let mut central = Vec::new();
42        for entry in &self.entries {
43            let crc = crc32(&entry.data);
44            let name = entry.name.as_bytes();
45            let size = entry.data.len() as u32;
46            file.write_all(b"PK\x03\x04")?;
47            file.write_all(&20u16.to_le_bytes())?;
48            file.write_all(&0u16.to_le_bytes())?;
49            file.write_all(&0u16.to_le_bytes())?;
50            file.write_all(&0u16.to_le_bytes())?;
51            file.write_all(&0x21u16.to_le_bytes())?;
52            file.write_all(&crc.to_le_bytes())?;
53            file.write_all(&size.to_le_bytes())?;
54            file.write_all(&size.to_le_bytes())?;
55            file.write_all(&(name.len() as u16).to_le_bytes())?;
56            file.write_all(&0u16.to_le_bytes())?;
57            file.write_all(name)?;
58            file.write_all(&entry.data)?;
59            central.extend_from_slice(b"PK\x01\x02");
60            central.extend_from_slice(&20u16.to_le_bytes());
61            central.extend_from_slice(&20u16.to_le_bytes());
62            central.extend_from_slice(&0u16.to_le_bytes());
63            central.extend_from_slice(&0u16.to_le_bytes());
64            central.extend_from_slice(&0u16.to_le_bytes());
65            central.extend_from_slice(&0x21u16.to_le_bytes());
66            central.extend_from_slice(&crc.to_le_bytes());
67            central.extend_from_slice(&size.to_le_bytes());
68            central.extend_from_slice(&size.to_le_bytes());
69            central.extend_from_slice(&(name.len() as u16).to_le_bytes());
70            central.extend_from_slice(&0u16.to_le_bytes());
71            central.extend_from_slice(&0u16.to_le_bytes());
72            central.extend_from_slice(&0u16.to_le_bytes());
73            central.extend_from_slice(&0u16.to_le_bytes());
74            central.extend_from_slice(&0u32.to_le_bytes());
75            central.extend_from_slice(&offset.to_le_bytes());
76            central.extend_from_slice(name);
77            offset += 30 + name.len() as u32 + size;
78        }
79        let central_size = central.len() as u32;
80        let central_offset = offset;
81        file.write_all(&central)?;
82        file.write_all(b"PK\x05\x06")?;
83        file.write_all(&0u16.to_le_bytes())?;
84        file.write_all(&0u16.to_le_bytes())?;
85        file.write_all(&(self.entries.len() as u16).to_le_bytes())?;
86        file.write_all(&(self.entries.len() as u16).to_le_bytes())?;
87        file.write_all(&central_size.to_le_bytes())?;
88        file.write_all(&central_offset.to_le_bytes())?;
89        file.write_all(&0u16.to_le_bytes())?;
90        file.flush()
91    }
92}
93
94impl Zip {
95    pub fn read<P: AsRef<Path>>(path: P) -> Result<Self> {
96        let mut file = BufReader::new(File::open(path)?);
97        let mut bytes = Vec::new();
98        file.read_to_end(&mut bytes)?;
99        Self::read_from(&bytes)
100    }
101
102    fn read_from(bytes: &[u8]) -> Result<Self> {
103        let eocd_sig = [0x50, 0x4b, 0x05, 0x06];
104        let pos = bytes
105            .windows(4)
106            .rposition(|w| w == eocd_sig)
107            .ok_or_else(|| invalid("not a zip file (no end of central directory)".into()))?;
108        if pos + 22 > bytes.len() {
109            return Err(invalid("truncated end of central directory".into()));
110        }
111        let eocd = &bytes[pos..pos + 22];
112        let total_entries = u16::from_le_bytes([eocd[10], eocd[11]]) as usize;
113        let central_size = u32::from_le_bytes([eocd[12], eocd[13], eocd[14], eocd[15]]) as usize;
114        let central_offset = u32::from_le_bytes([eocd[16], eocd[17], eocd[18], eocd[19]]) as usize;
115        if central_offset > bytes.len() || central_offset + central_size > bytes.len() {
116            return Err(invalid("central directory out of bounds".into()));
117        }
118        let mut entries = Vec::with_capacity(total_entries);
119        let mut cursor = central_offset;
120        for _ in 0..total_entries {
121            if cursor + 46 > bytes.len() || &bytes[cursor..cursor + 4] != b"PK\x01\x02" {
122                return Err(invalid("malformed central directory record".into()));
123            }
124            let record = &bytes[cursor..cursor + 46];
125            let method = u16::from_le_bytes([record[10], record[11]]);
126            let size =
127                u32::from_le_bytes([record[24], record[25], record[26], record[27]]) as usize;
128            let name_len = u16::from_le_bytes([record[28], record[29]]) as usize;
129            let extra_len = u16::from_le_bytes([record[30], record[31]]) as usize;
130            let comment_len = u16::from_le_bytes([record[32], record[33]]) as usize;
131            let local_offset =
132                u32::from_le_bytes([record[42], record[43], record[44], record[45]]) as usize;
133            let name_start = cursor + 46;
134            let name_end = name_start + name_len;
135            if name_end > bytes.len() {
136                return Err(invalid("truncated file name".into()));
137            }
138            let name = from_utf8(&bytes[name_start..name_end])
139                .map_err(|_| invalid("non-UTF-8 entry name".into()))?
140                .to_string();
141            if method != 0 {
142                return Err(invalid(format!(
143                    "unsupported compression method {method} for {name} (only stored entries are supported)"
144                )));
145            }
146            let data = read_local_entry(bytes, local_offset, size, &name)?;
147            entries.push(ZipEntry { name, data });
148            cursor = name_end + extra_len + comment_len;
149        }
150        Ok(Zip { entries })
151    }
152
153    pub fn entry(&self, name: &str) -> Option<&[u8]> {
154        self.entries
155            .iter()
156            .find(|entry| entry.name == name)
157            .map(|entry| entry.data.as_slice())
158    }
159}
160
161fn read_local_entry(bytes: &[u8], offset: usize, size: usize, name: &str) -> Result<Vec<u8>> {
162    if offset + 30 > bytes.len() || &bytes[offset..offset + 4] != b"PK\x03\x04" {
163        return Err(invalid(format!("malformed local file header for {name}")));
164    }
165    let name_len = u16::from_le_bytes([bytes[offset + 26], bytes[offset + 27]]) as usize;
166    let extra_len = u16::from_le_bytes([bytes[offset + 28], bytes[offset + 29]]) as usize;
167    let data_start = offset + 30 + name_len + extra_len;
168    let data_end = data_start + size;
169    if data_end > bytes.len() {
170        return Err(invalid(format!("truncated entry data for {name}")));
171    }
172    Ok(bytes[data_start..data_end].to_vec())
173}
174
175fn crc32(data: &[u8]) -> u32 {
176    let mut crc: u32 = 0xFFFF_FFFF;
177    for &byte in data {
178        crc ^= byte as u32;
179        for _ in 0..8 {
180            let mask = (crc & 1).wrapping_neg();
181            crc = (crc >> 1) ^ (0xEDB8_8320 & mask);
182        }
183    }
184    !crc
185}
186
187fn invalid(message: String) -> Error {
188    Error::new(ErrorKind::InvalidData, message)
189}