Skip to main content

zsync_rs/
control.rs

1use std::io::{BufRead, Read, Seek, SeekFrom, Write};
2
3use crate::checksum::{calc_md4, calc_sha1_stream};
4use crate::rsum::{Rsum, calc_rsum_block};
5
6#[derive(Debug, thiserror::Error)]
7pub enum GenerateError {
8    #[error("IO error: {0}")]
9    Io(#[from] std::io::Error),
10    #[error("file is empty")]
11    EmptyFile,
12}
13
14#[derive(Debug, thiserror::Error)]
15pub enum WriteError {
16    #[error("IO error: {0}")]
17    Io(#[from] std::io::Error),
18}
19
20#[derive(Debug, thiserror::Error)]
21pub enum ParseError {
22    #[error("IO error: {0}")]
23    Io(#[from] std::io::Error),
24    #[error("Invalid header: {0}")]
25    InvalidHeader(String),
26    #[error("Missing required field: {0}")]
27    MissingField(String),
28    #[error("Invalid blocksize: {0}")]
29    InvalidBlocksize(String),
30    #[error("Invalid hash lengths: {0}")]
31    InvalidHashLengths(String),
32    #[error("Invalid length: {0}")]
33    InvalidLength(String),
34    #[error("Unexpected end of file")]
35    UnexpectedEof,
36}
37
38#[derive(Debug, Clone, Copy)]
39pub struct BlockChecksum {
40    pub rsum: Rsum,
41    pub checksum: [u8; 16],
42}
43
44#[derive(Debug, Clone)]
45pub struct ControlFile {
46    pub version: String,
47    pub filename: Option<String>,
48    pub mtime: Option<String>,
49    pub blocksize: usize,
50    pub length: u64,
51    pub hash_lengths: HashLengths,
52    pub urls: Vec<String>,
53    pub sha1: Option<String>,
54    pub block_checksums: Vec<BlockChecksum>,
55}
56
57#[derive(Debug, Clone, Copy)]
58pub struct HashLengths {
59    pub seq_matches: u8,
60    pub rsum_bytes: u8,
61    pub checksum_bytes: u8,
62}
63
64impl Default for HashLengths {
65    fn default() -> Self {
66        Self {
67            seq_matches: 1,
68            rsum_bytes: 4,
69            checksum_bytes: 16,
70        }
71    }
72}
73
74impl ControlFile {
75    pub fn parse<R: Read>(reader: R) -> Result<Self, ParseError> {
76        let mut reader = std::io::BufReader::new(reader);
77        let mut line = String::new();
78
79        let mut version = String::new();
80        let mut filename = None;
81        let mut mtime = None;
82        let mut blocksize = None;
83        let mut length = None;
84        let mut hash_lengths = HashLengths::default();
85        let mut urls = Vec::new();
86        let mut sha1 = None;
87
88        loop {
89            line.clear();
90            let bytes_read = reader.read_line(&mut line)?;
91            if bytes_read == 0 {
92                return Err(ParseError::UnexpectedEof);
93            }
94
95            let trimmed = line.trim_end_matches(['\n', '\r', ' ']);
96            if trimmed.is_empty() {
97                break;
98            }
99
100            let Some((key, value)) = trimmed.split_once(':') else {
101                return Err(ParseError::InvalidHeader(trimmed.to_string()));
102            };
103
104            let value = value.trim_start_matches(' ');
105
106            match key {
107                "zsync" => {
108                    version = value.to_string();
109                }
110                "Filename" => {
111                    filename = Some(value.to_string());
112                }
113                "MTime" => {
114                    mtime = Some(value.to_string());
115                }
116                "Blocksize" => {
117                    let bs: usize = value
118                        .parse()
119                        .map_err(|_| ParseError::InvalidBlocksize(value.to_string()))?;
120                    if bs == 0 || (bs & (bs - 1)) != 0 {
121                        return Err(ParseError::InvalidBlocksize(value.to_string()));
122                    }
123                    blocksize = Some(bs);
124                }
125                "Length" => {
126                    length = Some(
127                        value
128                            .parse()
129                            .map_err(|_| ParseError::InvalidLength(value.to_string()))?,
130                    );
131                }
132                "URL" => {
133                    urls.push(value.to_string());
134                }
135                "Hash-Lengths" => {
136                    let parts: Vec<&str> = value.split(',').collect();
137                    if parts.len() != 3 {
138                        return Err(ParseError::InvalidHashLengths(value.to_string()));
139                    }
140                    let seq_matches: u8 = parts[0]
141                        .parse()
142                        .map_err(|_| ParseError::InvalidHashLengths(value.to_string()))?;
143                    let rsum_bytes: u8 = parts[1]
144                        .parse()
145                        .map_err(|_| ParseError::InvalidHashLengths(value.to_string()))?;
146                    let checksum_bytes: u8 = parts[2]
147                        .parse()
148                        .map_err(|_| ParseError::InvalidHashLengths(value.to_string()))?;
149
150                    if !(1..=2).contains(&seq_matches)
151                        || !(1..=4).contains(&rsum_bytes)
152                        || !(3..=16).contains(&checksum_bytes)
153                    {
154                        return Err(ParseError::InvalidHashLengths(value.to_string()));
155                    }
156
157                    hash_lengths = HashLengths {
158                        seq_matches,
159                        rsum_bytes,
160                        checksum_bytes,
161                    };
162                }
163                "SHA-1" => {
164                    if value.len() != 40 {
165                        return Err(ParseError::InvalidHeader(
166                            "SHA-1 digest wrong length".to_string(),
167                        ));
168                    }
169                    sha1 = Some(value.to_string());
170                }
171                _ => {}
172            }
173        }
174
175        let blocksize =
176            blocksize.ok_or_else(|| ParseError::MissingField("Blocksize".to_string()))?;
177        let length: u64 = length.ok_or_else(|| ParseError::MissingField("Length".to_string()))?;
178
179        let num_blocks = length.div_ceil(blocksize as u64) as usize;
180
181        // Sanity check: avoid massive allocations from malformed input.
182        // Each block needs (rsum_bytes + checksum_bytes) of data following the header.
183        const MAX_BLOCKS: usize = 64 * 1024 * 1024;
184        if num_blocks > MAX_BLOCKS {
185            return Err(ParseError::InvalidLength(format!(
186                "too many blocks: {num_blocks}"
187            )));
188        }
189
190        let block_checksums = Self::read_block_checksums(&mut reader, num_blocks, hash_lengths)?;
191
192        Ok(Self {
193            version,
194            filename,
195            mtime,
196            blocksize,
197            length,
198            hash_lengths,
199            urls,
200            sha1,
201            block_checksums,
202        })
203    }
204
205    fn read_block_checksums<R: BufRead>(
206        reader: &mut R,
207        num_blocks: usize,
208        hash_lengths: HashLengths,
209    ) -> Result<Vec<BlockChecksum>, ParseError> {
210        let mut checksums = Vec::with_capacity(num_blocks);
211        let entry_size = (hash_lengths.rsum_bytes + hash_lengths.checksum_bytes) as usize;
212        let mut buf = vec![0u8; entry_size];
213
214        for _ in 0..num_blocks {
215            reader.read_exact(&mut buf)?;
216
217            let rsum_bytes = hash_lengths.rsum_bytes as usize;
218            let (rsum_a, rsum_b) = match rsum_bytes {
219                1 => (0u16, u16::from(buf[0])),
220                2 => (0u16, u16::from_be_bytes([buf[0], buf[1]])),
221                3 => (u16::from(buf[0]), u16::from_be_bytes([buf[1], buf[2]])),
222                4 => (
223                    u16::from_be_bytes([buf[0], buf[1]]),
224                    u16::from_be_bytes([buf[2], buf[3]]),
225                ),
226                _ => (0, 0),
227            };
228
229            let mut checksum = [0u8; 16];
230            checksum[..hash_lengths.checksum_bytes as usize]
231                .copy_from_slice(&buf[rsum_bytes..entry_size]);
232
233            checksums.push(BlockChecksum {
234                rsum: Rsum {
235                    a: rsum_a,
236                    b: rsum_b,
237                },
238                checksum,
239            });
240        }
241
242        Ok(checksums)
243    }
244
245    pub fn num_blocks(&self) -> usize {
246        self.block_checksums.len()
247    }
248
249    /// Generate a control file by scanning an input file.
250    /// Blocksize is auto-calculated if `None`.
251    pub fn generate<R: Read + Seek>(
252        reader: &mut R,
253        filename: &str,
254        url: &str,
255        blocksize: Option<usize>,
256    ) -> Result<Self, GenerateError> {
257        let file_length = reader.seek(SeekFrom::End(0))?;
258        if file_length == 0 {
259            return Err(GenerateError::EmptyFile);
260        }
261        reader.seek(SeekFrom::Start(0))?;
262
263        let blocksize = blocksize.unwrap_or_else(|| auto_blocksize(file_length));
264        let hash_lengths = calculate_hash_lengths(file_length, blocksize);
265        let num_blocks = file_length.div_ceil(blocksize as u64) as usize;
266
267        let mut block_checksums = Vec::with_capacity(num_blocks);
268        let mut buf = vec![0u8; blocksize];
269
270        for i in 0..num_blocks {
271            let is_last = i == num_blocks - 1;
272            let block_len = if is_last {
273                let rem = (file_length % blocksize as u64) as usize;
274                if rem == 0 { blocksize } else { rem }
275            } else {
276                blocksize
277            };
278
279            reader.read_exact(&mut buf[..block_len])?;
280            if block_len < blocksize {
281                buf[block_len..].fill(0);
282            }
283
284            let rsum = calc_rsum_block(&buf);
285            let checksum = calc_md4(&buf);
286            block_checksums.push(BlockChecksum { rsum, checksum });
287        }
288
289        reader.seek(SeekFrom::Start(0))?;
290        let sha1_bytes = calc_sha1_stream(reader)?;
291        let sha1 = sha1_bytes
292            .iter()
293            .fold(String::with_capacity(40), |mut s, b| {
294                use std::fmt::Write;
295                let _ = write!(s, "{b:02x}");
296                s
297            });
298
299        Ok(Self {
300            version: "0.6.2".to_string(),
301            filename: Some(filename.to_string()),
302            mtime: None,
303            blocksize,
304            length: file_length,
305            hash_lengths,
306            urls: vec![url.to_string()],
307            sha1: Some(sha1),
308            block_checksums,
309        })
310    }
311
312    /// Write the control file to a writer.
313    pub fn write<W: Write>(&self, writer: &mut W) -> Result<(), WriteError> {
314        writeln!(writer, "zsync: {}", self.version)?;
315        if let Some(ref filename) = self.filename {
316            writeln!(writer, "Filename: {filename}")?;
317        }
318        if let Some(ref mtime) = self.mtime {
319            writeln!(writer, "MTime: {mtime}")?;
320        }
321        writeln!(writer, "Blocksize: {}", self.blocksize)?;
322        writeln!(writer, "Length: {}", self.length)?;
323        writeln!(
324            writer,
325            "Hash-Lengths: {},{},{}",
326            self.hash_lengths.seq_matches,
327            self.hash_lengths.rsum_bytes,
328            self.hash_lengths.checksum_bytes
329        )?;
330        for url in &self.urls {
331            writeln!(writer, "URL: {url}")?;
332        }
333        if let Some(ref sha1) = self.sha1 {
334            writeln!(writer, "SHA-1: {sha1}")?;
335        }
336        writeln!(writer)?;
337
338        let rsum_bytes = self.hash_lengths.rsum_bytes as usize;
339        let checksum_bytes = self.hash_lengths.checksum_bytes as usize;
340
341        for block in &self.block_checksums {
342            let rsum_be = rsum_to_bytes(block.rsum, rsum_bytes);
343            writer.write_all(&rsum_be)?;
344            writer.write_all(&block.checksum[..checksum_bytes])?;
345        }
346
347        Ok(())
348    }
349}
350
351fn rsum_to_bytes(rsum: Rsum, rsum_bytes: usize) -> Vec<u8> {
352    match rsum_bytes {
353        1 => vec![rsum.b as u8],
354        2 => rsum.b.to_be_bytes().to_vec(),
355        3 => {
356            let mut v = Vec::with_capacity(3);
357            v.push(rsum.a as u8);
358            v.extend_from_slice(&rsum.b.to_be_bytes());
359            v
360        }
361        4 => {
362            let mut v = Vec::with_capacity(4);
363            v.extend_from_slice(&rsum.a.to_be_bytes());
364            v.extend_from_slice(&rsum.b.to_be_bytes());
365            v
366        }
367        _ => vec![0; rsum_bytes],
368    }
369}
370
371fn auto_blocksize(file_length: u64) -> usize {
372    if file_length < 100_000_000 {
373        2048
374    } else {
375        4096
376    }
377}
378
379/// Calculate optimal hash lengths based on file size and blocksize.
380fn calculate_hash_lengths(file_length: u64, blocksize: usize) -> HashLengths {
381    let len = file_length as f64;
382    let bs = blocksize as f64;
383    let seq_matches: u8 = if file_length > blocksize as u64 { 2 } else { 1 };
384    let sm = f64::from(seq_matches);
385
386    let rsum_bytes = ((len.ln() + bs.ln()) / 2_f64.ln() - 8.6) / sm / 8.0;
387    let rsum_bytes = (rsum_bytes.ceil() as i32).clamp(2, 4);
388
389    let num_blocks = 1.0 + len / bs;
390    let calc1 = ((20.0 + len.log2() + num_blocks.log2()) / sm / 8.0).ceil();
391    let calc2 = (7.9 + (20.0 + num_blocks.log2())) / 8.0;
392    let checksum_bytes = (calc1.max(calc2) as i32).clamp(4, 16);
393
394    HashLengths {
395        seq_matches,
396        rsum_bytes: rsum_bytes as u8,
397        checksum_bytes: checksum_bytes as u8,
398    }
399}
400
401#[cfg(test)]
402mod tests {
403    use super::*;
404
405    #[test]
406    fn test_parse_minimal() {
407        let mut data = Vec::new();
408        data.extend_from_slice(
409            b"zsync: 0.6.2\nBlocksize: 2048\nLength: 2048\nHash-Lengths: 1,4,16\n\n",
410        );
411        data.extend_from_slice(&[0u8; 20]);
412        let result = ControlFile::parse(&data[..]);
413        assert!(result.is_ok());
414        let cf = result.unwrap();
415        assert_eq!(cf.blocksize, 2048);
416        assert_eq!(cf.length, 2048);
417    }
418
419    #[test]
420    fn test_parse_missing_blocksize() {
421        let data = b"zsync: 0.6.2\nLength: 4096\n\n";
422        let result = ControlFile::parse(&data[..]);
423        assert!(result.is_err());
424    }
425
426    #[test]
427    fn test_parse_invalid_blocksize() {
428        let data = b"zsync: 0.6.2\nBlocksize: 1000\nLength: 4096\n\n";
429        let result = ControlFile::parse(&data[..]);
430        assert!(result.is_err());
431    }
432
433    #[test]
434    fn test_generate_write_roundtrip() {
435        let file_data = vec![42u8; 4096];
436        let mut cursor = std::io::Cursor::new(&file_data);
437
438        let cf = ControlFile::generate(&mut cursor, "test.bin", "test.bin", Some(2048)).unwrap();
439        assert_eq!(cf.blocksize, 2048);
440        assert_eq!(cf.length, 4096);
441        assert_eq!(cf.block_checksums.len(), 2);
442        assert!(cf.sha1.is_some());
443
444        let mut buf = Vec::new();
445        cf.write(&mut buf).unwrap();
446        let parsed = ControlFile::parse(&buf[..]).unwrap();
447
448        assert_eq!(parsed.blocksize, cf.blocksize);
449        assert_eq!(parsed.length, cf.length);
450        assert_eq!(parsed.sha1, cf.sha1);
451        assert_eq!(parsed.block_checksums.len(), cf.block_checksums.len());
452        let rlen = cf.hash_lengths.rsum_bytes as usize;
453        let clen = cf.hash_lengths.checksum_bytes as usize;
454        for (a, b) in parsed.block_checksums.iter().zip(&cf.block_checksums) {
455            let a_bytes = rsum_to_bytes(a.rsum, rlen);
456            let b_bytes = rsum_to_bytes(b.rsum, rlen);
457            assert_eq!(a_bytes, b_bytes);
458            assert_eq!(a.checksum[..clen], b.checksum[..clen]);
459        }
460    }
461
462    #[test]
463    fn test_generate_empty_file() {
464        let file_data: Vec<u8> = vec![];
465        let mut cursor = std::io::Cursor::new(&file_data);
466        let result = ControlFile::generate(&mut cursor, "empty", "empty", None);
467        assert!(result.is_err());
468    }
469
470    #[test]
471    fn test_generate_partial_last_block() {
472        // File not aligned to blocksize
473        let file_data = vec![0xABu8; 3000];
474        let mut cursor = std::io::Cursor::new(&file_data);
475
476        let cf = ControlFile::generate(&mut cursor, "test.bin", "test.bin", Some(2048)).unwrap();
477        assert_eq!(cf.length, 3000);
478        assert_eq!(cf.block_checksums.len(), 2);
479    }
480
481    #[test]
482    fn test_auto_blocksize_small() {
483        assert_eq!(auto_blocksize(1024), 2048);
484        assert_eq!(auto_blocksize(99_999_999), 2048);
485    }
486
487    #[test]
488    fn test_auto_blocksize_large() {
489        assert_eq!(auto_blocksize(100_000_000), 4096);
490        assert_eq!(auto_blocksize(500_000_000), 4096);
491    }
492
493    #[test]
494    fn test_rsum_to_bytes() {
495        let rsum = Rsum {
496            a: 0x1234,
497            b: 0x5678,
498        };
499        assert_eq!(rsum_to_bytes(rsum, 4), vec![0x12, 0x34, 0x56, 0x78]);
500        assert_eq!(rsum_to_bytes(rsum, 2), vec![0x56, 0x78]);
501        assert_eq!(rsum_to_bytes(rsum, 1), vec![0x78]);
502    }
503}