1use crate::codec::compress::Algorithm;
13use crate::error::{Error, Result};
14use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
15
16pub const WIRE_VERSION: u16 = 1;
17const FRAME_MAGIC: u8 = 0x52; pub const FRAME_HEADER_LEN: usize = 28;
22
23pub mod flags {
24 pub const LAST_CHUNK: u8 = 1 << 0;
26 pub const SEALED: u8 = 1 << 1;
28 pub const ZERO: u8 = 1 << 2;
32 pub const REUSE_LOCAL: u8 = 1 << 3;
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq)]
39pub struct FrameHeader {
40 pub flags: u8,
41 pub algorithm: Algorithm,
42 pub file_id: u32,
43 pub chunk_index: u64,
44 pub epoch: u32,
46 pub raw_len: u32,
48 pub payload_len: u32,
50}
51
52impl FrameHeader {
53 pub fn encode(&self, out: &mut [u8; FRAME_HEADER_LEN]) {
54 out[0] = FRAME_MAGIC;
55 out[1] = WIRE_VERSION as u8;
56 out[2] = self.flags;
57 out[3] = self.algorithm as u8;
58 out[4..8].copy_from_slice(&self.file_id.to_le_bytes());
59 out[8..16].copy_from_slice(&self.chunk_index.to_le_bytes());
60 out[16..20].copy_from_slice(&self.epoch.to_le_bytes());
61 out[20..24].copy_from_slice(&self.raw_len.to_le_bytes());
62 out[24..28].copy_from_slice(&self.payload_len.to_le_bytes());
63 }
64
65 pub fn decode(buf: &[u8; FRAME_HEADER_LEN], max_frame: usize) -> Result<Self> {
66 if buf[0] != FRAME_MAGIC {
67 return Err(Error::protocol(format!("bad frame magic 0x{:02x}", buf[0])));
68 }
69 if buf[1] != WIRE_VERSION as u8 {
70 return Err(Error::Version {
71 peer: buf[1] as u16,
72 ours: WIRE_VERSION,
73 });
74 }
75 let payload_len = u32::from_le_bytes(buf[24..28].try_into().unwrap());
76 let raw_len = u32::from_le_bytes(buf[20..24].try_into().unwrap());
77 if payload_len as usize > max_frame {
79 return Err(Error::FrameTooLarge {
80 got: payload_len as usize,
81 limit: max_frame,
82 });
83 }
84 if raw_len as usize > max_frame {
85 return Err(Error::FrameTooLarge {
86 got: raw_len as usize,
87 limit: max_frame,
88 });
89 }
90 Ok(FrameHeader {
91 flags: buf[2],
92 algorithm: Algorithm::from_u8(buf[3])?,
93 file_id: u32::from_le_bytes(buf[4..8].try_into().unwrap()),
94 chunk_index: u64::from_le_bytes(buf[8..16].try_into().unwrap()),
95 epoch: u32::from_le_bytes(buf[16..20].try_into().unwrap()),
96 raw_len,
97 payload_len,
98 })
99 }
100
101 pub fn sealed(&self) -> bool {
102 self.flags & flags::SEALED != 0
103 }
104
105 pub fn last_chunk(&self) -> bool {
106 self.flags & flags::LAST_CHUNK != 0
107 }
108
109 pub fn is_zero(&self) -> bool {
110 self.flags & flags::ZERO != 0
111 }
112
113 pub fn reuses_local(&self) -> bool {
114 self.flags & flags::REUSE_LOCAL != 0
115 }
116
117 pub fn wire_payload_len(&self) -> usize {
119 self.payload_len as usize
120 + if self.sealed() {
121 crate::codec::crypto::TAG_LEN
122 } else {
123 0
124 }
125 }
126}
127
128#[derive(Debug, Clone, Copy, PartialEq, Eq)]
133#[repr(u8)]
134pub enum ControlKind {
135 Manifest = 1,
136 ResumeState = 2,
137 Start = 3,
138 LocalIndex = 7,
139 FileComplete = 4,
140 AllComplete = 5,
141 Abort = 6,
142}
143
144impl ControlKind {
145 fn from_u8(v: u8) -> Result<Self> {
146 Ok(match v {
147 1 => ControlKind::Manifest,
148 2 => ControlKind::ResumeState,
149 3 => ControlKind::Start,
150 4 => ControlKind::FileComplete,
151 5 => ControlKind::AllComplete,
152 6 => ControlKind::Abort,
153 7 => ControlKind::LocalIndex,
154 other => return Err(Error::protocol(format!("unknown control message {other}"))),
155 })
156 }
157}
158
159#[derive(Debug, Clone, PartialEq, Eq)]
161pub struct FileEntry {
162 pub file_id: u32,
163 pub path: String,
165 pub size: u64,
166 pub chunk_size: u32,
167 pub mode: u32,
169 pub mtime: i64,
171 pub kind: EntryKind,
172 pub hash: Option<[u8; 32]>,
174 pub incompressible: bool,
176}
177
178#[derive(Debug, Clone, Copy, PartialEq, Eq)]
179#[repr(u8)]
180pub enum EntryKind {
181 File = 0,
182 Directory = 1,
183 Symlink = 2,
184}
185
186impl EntryKind {
187 fn from_u8(v: u8) -> Result<Self> {
188 Ok(match v {
189 0 => EntryKind::File,
190 1 => EntryKind::Directory,
191 2 => EntryKind::Symlink,
192 other => return Err(Error::protocol(format!("unknown entry kind {other}"))),
193 })
194 }
195}
196
197impl FileEntry {
198 pub fn chunk_count(&self) -> u64 {
199 if self.chunk_size == 0 {
200 return 0;
201 }
202 self.size.div_ceil(self.chunk_size as u64)
203 }
204
205 fn encode(&self, out: &mut Vec<u8>) {
206 out.extend_from_slice(&self.file_id.to_le_bytes());
207 out.extend_from_slice(&self.size.to_le_bytes());
208 out.extend_from_slice(&self.chunk_size.to_le_bytes());
209 out.extend_from_slice(&self.mode.to_le_bytes());
210 out.extend_from_slice(&self.mtime.to_le_bytes());
211 out.push(self.kind as u8);
212 let mut bits = 0u8;
213 if self.hash.is_some() {
214 bits |= 1;
215 }
216 if self.incompressible {
217 bits |= 2;
218 }
219 out.push(bits);
220 let p = self.path.as_bytes();
221 out.extend_from_slice(&(p.len() as u16).to_le_bytes());
222 out.extend_from_slice(p);
223 if let Some(h) = &self.hash {
224 out.extend_from_slice(h);
225 }
226 }
227
228 fn decode(cur: &mut Cursor<'_>) -> Result<Self> {
229 let file_id = cur.u32()?;
230 let size = cur.u64()?;
231 let chunk_size = cur.u32()?;
232 let mode = cur.u32()?;
233 let mtime = cur.i64()?;
234 let kind = EntryKind::from_u8(cur.u8()?)?;
235 let bits = cur.u8()?;
236 let path_len = cur.u16()? as usize;
237 if path_len > 4096 {
240 return Err(Error::protocol("manifest path exceeds 4096 bytes"));
241 }
242 let path_bytes = cur.take(path_len)?;
243 let path = String::from_utf8(path_bytes.to_vec())
244 .map_err(|_| Error::protocol("manifest path is not valid UTF-8"))?;
245 let hash = if bits & 1 != 0 {
246 let h = cur.take(32)?;
247 let mut a = [0u8; 32];
248 a.copy_from_slice(h);
249 Some(a)
250 } else {
251 None
252 };
253 if chunk_size == 0 && kind == EntryKind::File && size > 0 {
254 return Err(Error::protocol("non-empty file entry has chunk_size 0"));
255 }
256 Ok(FileEntry {
257 file_id,
258 path,
259 size,
260 chunk_size,
261 mode,
262 mtime,
263 kind,
264 hash,
265 incompressible: bits & 2 != 0,
266 })
267 }
268}
269
270#[derive(Debug, Clone, PartialEq, Eq)]
273pub struct LocalFileIndex {
274 pub file_id: u32,
275 pub hashes: Vec<[u8; 32]>,
278}
279
280#[derive(Debug, Clone, PartialEq, Eq)]
281pub struct ResumeEntry {
282 pub file_id: u32,
283 pub have: Vec<u8>,
285}
286
287#[derive(Debug, Clone, PartialEq, Eq)]
288pub enum Control {
289 Manifest(Vec<FileEntry>),
290 ResumeState(Vec<ResumeEntry>),
291 LocalIndex(Vec<LocalFileIndex>),
300 Start {
303 streams: u32,
304 },
305 FileComplete {
308 file_id: u32,
309 hash: Option<[u8; 32]>,
310 },
311 AllComplete,
312 Abort {
313 reason: String,
314 },
315}
316
317impl Control {
318 fn kind(&self) -> ControlKind {
319 match self {
320 Control::LocalIndex(_) => ControlKind::LocalIndex,
321 Control::Manifest(_) => ControlKind::Manifest,
322 Control::ResumeState(_) => ControlKind::ResumeState,
323 Control::Start { .. } => ControlKind::Start,
324 Control::FileComplete { .. } => ControlKind::FileComplete,
325 Control::AllComplete => ControlKind::AllComplete,
326 Control::Abort { .. } => ControlKind::Abort,
327 }
328 }
329
330 fn encode_body(&self, out: &mut Vec<u8>) {
331 match self {
332 Control::Manifest(entries) => {
333 out.extend_from_slice(&(entries.len() as u32).to_le_bytes());
334 for e in entries {
335 e.encode(out);
336 }
337 }
338 Control::ResumeState(entries) => {
339 out.extend_from_slice(&(entries.len() as u32).to_le_bytes());
340 for e in entries {
341 out.extend_from_slice(&e.file_id.to_le_bytes());
342 out.extend_from_slice(&(e.have.len() as u32).to_le_bytes());
343 out.extend_from_slice(&e.have);
344 }
345 }
346 Control::LocalIndex(entries) => {
347 out.extend_from_slice(&(entries.len() as u32).to_le_bytes());
348 for e in entries {
349 out.extend_from_slice(&e.file_id.to_le_bytes());
350 out.extend_from_slice(&(e.hashes.len() as u32).to_le_bytes());
351 for h in &e.hashes {
352 out.extend_from_slice(h);
353 }
354 }
355 }
356 Control::AllComplete => {}
357 Control::Start { streams } => {
358 out.extend_from_slice(&streams.to_le_bytes());
359 }
360 Control::FileComplete { file_id, hash } => {
361 out.extend_from_slice(&file_id.to_le_bytes());
362 match hash {
363 Some(h) => {
364 out.push(1);
365 out.extend_from_slice(h);
366 }
367 None => out.push(0),
368 }
369 }
370 Control::Abort { reason } => {
371 let r = reason.as_bytes();
372 let n = r.len().min(1024);
373 out.extend_from_slice(&(n as u16).to_le_bytes());
374 out.extend_from_slice(&r[..n]);
375 }
376 }
377 }
378
379 fn decode_body(kind: ControlKind, body: &[u8], max_entries: usize) -> Result<Self> {
380 let mut cur = Cursor::new(body);
381 Ok(match kind {
382 ControlKind::Manifest => {
383 let n = cur.u32()? as usize;
384 if n > max_entries {
385 return Err(Error::protocol(format!(
386 "manifest declares {n} entries, limit is {max_entries}"
387 )));
388 }
389 let mut v = Vec::with_capacity(n.min(4096));
390 for _ in 0..n {
391 v.push(FileEntry::decode(&mut cur)?);
392 }
393 Control::Manifest(v)
394 }
395 ControlKind::ResumeState => {
396 let n = cur.u32()? as usize;
397 if n > max_entries {
398 return Err(Error::protocol("resume state entry count exceeds limit"));
399 }
400 let mut v = Vec::with_capacity(n.min(4096));
401 for _ in 0..n {
402 let file_id = cur.u32()?;
403 let len = cur.u32()? as usize;
404 let have = cur.take(len)?.to_vec();
405 v.push(ResumeEntry { file_id, have });
406 }
407 Control::ResumeState(v)
408 }
409 ControlKind::LocalIndex => {
410 let n = cur.u32()? as usize;
411 if n > max_entries {
412 return Err(Error::protocol("local index entry count exceeds limit"));
413 }
414 let mut v = Vec::with_capacity(n.min(4096));
415 for _ in 0..n {
416 let file_id = cur.u32()?;
417 let count = cur.u32()? as usize;
418 if count > (1 << 26) {
421 return Err(Error::protocol("local index declares too many chunks"));
422 }
423 let mut hashes = Vec::with_capacity(count.min(4096));
424 for _ in 0..count {
425 let mut h = [0u8; 32];
426 h.copy_from_slice(cur.take(32)?);
427 hashes.push(h);
428 }
429 v.push(LocalFileIndex { file_id, hashes });
430 }
431 Control::LocalIndex(v)
432 }
433 ControlKind::Start => Control::Start {
434 streams: cur.u32()?,
435 },
436 ControlKind::AllComplete => Control::AllComplete,
437 ControlKind::FileComplete => {
438 let file_id = cur.u32()?;
439 let hash = if cur.u8()? == 1 {
440 let mut h = [0u8; 32];
441 h.copy_from_slice(cur.take(32)?);
442 Some(h)
443 } else {
444 None
445 };
446 Control::FileComplete { file_id, hash }
447 }
448 ControlKind::Abort => {
449 let n = cur.u16()? as usize;
450 let s = cur.take(n)?;
451 Control::Abort {
452 reason: String::from_utf8_lossy(s).into_owned(),
453 }
454 }
455 })
456 }
457}
458
459pub async fn write_control<W: AsyncWrite + Unpin>(w: &mut W, msg: &Control) -> Result<()> {
461 let mut body = Vec::with_capacity(64);
462 msg.encode_body(&mut body);
463 let mut framed = Vec::with_capacity(body.len() + 5);
464 framed.push(msg.kind() as u8);
465 framed.extend_from_slice(&(body.len() as u32).to_le_bytes());
466 framed.extend_from_slice(&body);
467 w.write_all(&framed).await?;
468 w.flush().await?;
469 Ok(())
470}
471
472pub async fn read_control<R: AsyncRead + Unpin>(
474 r: &mut R,
475 max_body: usize,
476 max_entries: usize,
477) -> Result<Control> {
478 let mut hdr = [0u8; 5];
479 r.read_exact(&mut hdr).await.map_err(map_eof)?;
480 let kind = ControlKind::from_u8(hdr[0])?;
481 let len = u32::from_le_bytes(hdr[1..5].try_into().unwrap()) as usize;
482 if len > max_body {
483 return Err(Error::FrameTooLarge {
484 got: len,
485 limit: max_body,
486 });
487 }
488 let mut body = vec![0u8; len];
489 r.read_exact(&mut body).await.map_err(map_eof)?;
490 Control::decode_body(kind, &body, max_entries)
491}
492
493fn map_eof(e: std::io::Error) -> Error {
494 if e.kind() == std::io::ErrorKind::UnexpectedEof {
495 Error::Closed("stream ended mid-message".into())
496 } else {
497 Error::Io(e)
498 }
499}
500
501struct Cursor<'a> {
509 buf: &'a [u8],
510 pos: usize,
511}
512
513impl<'a> Cursor<'a> {
514 fn new(buf: &'a [u8]) -> Self {
515 Self { buf, pos: 0 }
516 }
517 fn take(&mut self, n: usize) -> Result<&'a [u8]> {
518 let end = self
519 .pos
520 .checked_add(n)
521 .ok_or_else(|| Error::protocol("length overflow"))?;
522 if end > self.buf.len() {
523 return Err(Error::protocol("message truncated"));
524 }
525 let s = &self.buf[self.pos..end];
526 self.pos = end;
527 Ok(s)
528 }
529 fn u8(&mut self) -> Result<u8> {
530 Ok(self.take(1)?[0])
531 }
532 fn u16(&mut self) -> Result<u16> {
533 Ok(u16::from_le_bytes(self.take(2)?.try_into().unwrap()))
534 }
535 fn u32(&mut self) -> Result<u32> {
536 Ok(u32::from_le_bytes(self.take(4)?.try_into().unwrap()))
537 }
538 fn u64(&mut self) -> Result<u64> {
539 Ok(u64::from_le_bytes(self.take(8)?.try_into().unwrap()))
540 }
541 fn i64(&mut self) -> Result<i64> {
542 Ok(i64::from_le_bytes(self.take(8)?.try_into().unwrap()))
543 }
544}
545
546#[cfg(test)]
547mod tests {
548 use super::*;
549
550 fn sample_entry(id: u32) -> FileEntry {
551 FileEntry {
552 file_id: id,
553 path: format!("music/album {id}/track.flac"),
554 size: 987_654_321_000,
555 chunk_size: 1 << 20,
556 mode: 0o644,
557 mtime: 1_700_000_000,
558 kind: EntryKind::File,
559 hash: Some([id as u8; 32]),
560 incompressible: true,
561 }
562 }
563
564 #[test]
565 fn frame_header_roundtrip() {
566 let h = FrameHeader {
567 flags: flags::LAST_CHUNK | flags::SEALED,
568 algorithm: Algorithm::Zstd,
569 file_id: 0xDEAD_BEEF,
570 chunk_index: 9_876_543_210,
572 epoch: 3,
573 raw_len: 1 << 20,
574 payload_len: 700_000,
575 };
576 let mut buf = [0u8; FRAME_HEADER_LEN];
577 h.encode(&mut buf);
578 let d = FrameHeader::decode(&buf, 64 << 20).unwrap();
579 assert_eq!(h, d);
580 assert!(d.sealed() && d.last_chunk());
581 assert_eq!(d.wire_payload_len(), 700_000 + 16);
582 }
583
584 #[test]
585 fn frame_header_rejects_oversized_lengths() {
586 let h = FrameHeader {
587 flags: 0,
588 algorithm: Algorithm::None,
589 file_id: 1,
590 chunk_index: 0,
591 epoch: 0,
592 raw_len: 16,
593 payload_len: u32::MAX,
594 };
595 let mut buf = [0u8; FRAME_HEADER_LEN];
596 h.encode(&mut buf);
597 assert!(matches!(
598 FrameHeader::decode(&buf, 1 << 20),
599 Err(Error::FrameTooLarge { .. })
600 ));
601 }
602
603 #[test]
604 fn frame_header_rejects_bad_magic_and_version() {
605 let mut buf = [0u8; FRAME_HEADER_LEN];
606 FrameHeader {
607 flags: 0,
608 algorithm: Algorithm::None,
609 file_id: 1,
610 chunk_index: 0,
611 epoch: 0,
612 raw_len: 1,
613 payload_len: 1,
614 }
615 .encode(&mut buf);
616 let mut bad = buf;
617 bad[0] = 0x00;
618 assert!(FrameHeader::decode(&bad, 1 << 20).is_err());
619 let mut bad = buf;
620 bad[1] = 99;
621 assert!(matches!(
622 FrameHeader::decode(&bad, 1 << 20),
623 Err(Error::Version { .. })
624 ));
625 }
626
627 #[tokio::test]
628 async fn control_roundtrip_all_variants() {
629 let msgs = vec![
630 Control::Manifest((0..64).map(sample_entry).collect()),
631 Control::ResumeState(vec![ResumeEntry {
632 file_id: 3,
633 have: vec![0xFF, 0x0F, 0x00],
634 }]),
635 Control::LocalIndex(vec![LocalFileIndex {
636 file_id: 4,
637 hashes: vec![[1u8; 32], [2u8; 32], [3u8; 32]],
638 }]),
639 Control::Start { streams: 16 },
640 Control::FileComplete {
641 file_id: 12,
642 hash: Some([7u8; 32]),
643 },
644 Control::FileComplete {
645 file_id: 13,
646 hash: None,
647 },
648 Control::AllComplete,
649 Control::Abort {
650 reason: "disk full".into(),
651 },
652 ];
653 for m in msgs {
654 let mut buf = Vec::new();
655 write_control(&mut buf, &m).await.unwrap();
656 let mut slice = &buf[..];
657 let got = read_control(&mut slice, 16 << 20, 1 << 20).await.unwrap();
658 assert_eq!(m, got);
659 }
660 }
661
662 #[tokio::test]
663 async fn control_rejects_oversized_body() {
664 let m = Control::Manifest((0..1000).map(sample_entry).collect());
665 let mut buf = Vec::new();
666 write_control(&mut buf, &m).await.unwrap();
667 let mut slice = &buf[..];
668 assert!(matches!(
669 read_control(&mut slice, 128, 1 << 20).await,
670 Err(Error::FrameTooLarge { .. })
671 ));
672 }
673
674 #[tokio::test]
675 async fn control_rejects_absurd_entry_count() {
676 let mut body = Vec::new();
677 body.extend_from_slice(&(5_000_000u32).to_le_bytes());
678 let mut framed = vec![ControlKind::Manifest as u8];
679 framed.extend_from_slice(&(body.len() as u32).to_le_bytes());
680 framed.extend_from_slice(&body);
681 let mut slice = &framed[..];
682 assert!(read_control(&mut slice, 16 << 20, 1000).await.is_err());
684 }
685
686 #[tokio::test]
687 async fn truncated_manifest_is_an_error_not_a_panic() {
688 let m = Control::Manifest((0..8).map(sample_entry).collect());
689 let mut buf = Vec::new();
690 write_control(&mut buf, &m).await.unwrap();
691 let cut = buf.len() - 40;
693 let body_len = u32::from_le_bytes(buf[1..5].try_into().unwrap()) as usize;
694 let truncated_body = &buf[5..cut];
695 let r = Control::decode_body(ControlKind::Manifest, truncated_body, 1 << 20);
696 assert!(r.is_err());
697 assert!(body_len > truncated_body.len());
698 }
699
700 #[test]
701 fn chunk_count_is_exact_at_boundaries() {
702 let mut e = sample_entry(1);
703 e.chunk_size = 1024;
704 e.size = 0;
705 assert_eq!(e.chunk_count(), 0);
706 e.size = 1;
707 assert_eq!(e.chunk_count(), 1);
708 e.size = 1024;
709 assert_eq!(e.chunk_count(), 1);
710 e.size = 1025;
711 assert_eq!(e.chunk_count(), 2);
712 e.chunk_size = 1 << 20;
714 e.size = 100 * (1 << 30);
715 assert_eq!(e.chunk_count(), 102_400);
716 }
717}