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 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 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 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
379fn 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 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}