use crate::codec::compress::Algorithm;
use crate::error::{Error, Result};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
pub const WIRE_VERSION: u16 = 1;
const FRAME_MAGIC: u8 = 0x52;
pub const FRAME_HEADER_LEN: usize = 28;
pub mod flags {
pub const LAST_CHUNK: u8 = 1 << 0;
pub const SEALED: u8 = 1 << 1;
pub const ZERO: u8 = 1 << 2;
pub const REUSE_LOCAL: u8 = 1 << 3;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FrameHeader {
pub flags: u8,
pub algorithm: Algorithm,
pub file_id: u32,
pub chunk_index: u64,
pub epoch: u32,
pub raw_len: u32,
pub payload_len: u32,
}
impl FrameHeader {
pub fn encode(&self, out: &mut [u8; FRAME_HEADER_LEN]) {
out[0] = FRAME_MAGIC;
out[1] = WIRE_VERSION as u8;
out[2] = self.flags;
out[3] = self.algorithm as u8;
out[4..8].copy_from_slice(&self.file_id.to_le_bytes());
out[8..16].copy_from_slice(&self.chunk_index.to_le_bytes());
out[16..20].copy_from_slice(&self.epoch.to_le_bytes());
out[20..24].copy_from_slice(&self.raw_len.to_le_bytes());
out[24..28].copy_from_slice(&self.payload_len.to_le_bytes());
}
pub fn decode(buf: &[u8; FRAME_HEADER_LEN], max_frame: usize) -> Result<Self> {
if buf[0] != FRAME_MAGIC {
return Err(Error::protocol(format!("bad frame magic 0x{:02x}", buf[0])));
}
if buf[1] != WIRE_VERSION as u8 {
return Err(Error::Version {
peer: buf[1] as u16,
ours: WIRE_VERSION,
});
}
let payload_len = u32::from_le_bytes(buf[24..28].try_into().unwrap());
let raw_len = u32::from_le_bytes(buf[20..24].try_into().unwrap());
if payload_len as usize > max_frame {
return Err(Error::FrameTooLarge {
got: payload_len as usize,
limit: max_frame,
});
}
if raw_len as usize > max_frame {
return Err(Error::FrameTooLarge {
got: raw_len as usize,
limit: max_frame,
});
}
Ok(FrameHeader {
flags: buf[2],
algorithm: Algorithm::from_u8(buf[3])?,
file_id: u32::from_le_bytes(buf[4..8].try_into().unwrap()),
chunk_index: u64::from_le_bytes(buf[8..16].try_into().unwrap()),
epoch: u32::from_le_bytes(buf[16..20].try_into().unwrap()),
raw_len,
payload_len,
})
}
pub fn sealed(&self) -> bool {
self.flags & flags::SEALED != 0
}
pub fn last_chunk(&self) -> bool {
self.flags & flags::LAST_CHUNK != 0
}
pub fn is_zero(&self) -> bool {
self.flags & flags::ZERO != 0
}
pub fn reuses_local(&self) -> bool {
self.flags & flags::REUSE_LOCAL != 0
}
pub fn wire_payload_len(&self) -> usize {
self.payload_len as usize
+ if self.sealed() {
crate::codec::crypto::TAG_LEN
} else {
0
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum ControlKind {
Manifest = 1,
ResumeState = 2,
Start = 3,
LocalIndex = 7,
FileComplete = 4,
AllComplete = 5,
Abort = 6,
}
impl ControlKind {
fn from_u8(v: u8) -> Result<Self> {
Ok(match v {
1 => ControlKind::Manifest,
2 => ControlKind::ResumeState,
3 => ControlKind::Start,
4 => ControlKind::FileComplete,
5 => ControlKind::AllComplete,
6 => ControlKind::Abort,
7 => ControlKind::LocalIndex,
other => return Err(Error::protocol(format!("unknown control message {other}"))),
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FileEntry {
pub file_id: u32,
pub path: String,
pub size: u64,
pub chunk_size: u32,
pub mode: u32,
pub mtime: i64,
pub kind: EntryKind,
pub hash: Option<[u8; 32]>,
pub incompressible: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum EntryKind {
File = 0,
Directory = 1,
Symlink = 2,
}
impl EntryKind {
fn from_u8(v: u8) -> Result<Self> {
Ok(match v {
0 => EntryKind::File,
1 => EntryKind::Directory,
2 => EntryKind::Symlink,
other => return Err(Error::protocol(format!("unknown entry kind {other}"))),
})
}
}
impl FileEntry {
pub fn chunk_count(&self) -> u64 {
if self.chunk_size == 0 {
return 0;
}
self.size.div_ceil(self.chunk_size as u64)
}
fn encode(&self, out: &mut Vec<u8>) {
out.extend_from_slice(&self.file_id.to_le_bytes());
out.extend_from_slice(&self.size.to_le_bytes());
out.extend_from_slice(&self.chunk_size.to_le_bytes());
out.extend_from_slice(&self.mode.to_le_bytes());
out.extend_from_slice(&self.mtime.to_le_bytes());
out.push(self.kind as u8);
let mut bits = 0u8;
if self.hash.is_some() {
bits |= 1;
}
if self.incompressible {
bits |= 2;
}
out.push(bits);
let p = self.path.as_bytes();
out.extend_from_slice(&(p.len() as u16).to_le_bytes());
out.extend_from_slice(p);
if let Some(h) = &self.hash {
out.extend_from_slice(h);
}
}
fn decode(cur: &mut Cursor<'_>) -> Result<Self> {
let file_id = cur.u32()?;
let size = cur.u64()?;
let chunk_size = cur.u32()?;
let mode = cur.u32()?;
let mtime = cur.i64()?;
let kind = EntryKind::from_u8(cur.u8()?)?;
let bits = cur.u8()?;
let path_len = cur.u16()? as usize;
if path_len > 4096 {
return Err(Error::protocol("manifest path exceeds 4096 bytes"));
}
let path_bytes = cur.take(path_len)?;
let path = String::from_utf8(path_bytes.to_vec())
.map_err(|_| Error::protocol("manifest path is not valid UTF-8"))?;
let hash = if bits & 1 != 0 {
let h = cur.take(32)?;
let mut a = [0u8; 32];
a.copy_from_slice(h);
Some(a)
} else {
None
};
if chunk_size == 0 && kind == EntryKind::File && size > 0 {
return Err(Error::protocol("non-empty file entry has chunk_size 0"));
}
Ok(FileEntry {
file_id,
path,
size,
chunk_size,
mode,
mtime,
kind,
hash,
incompressible: bits & 2 != 0,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LocalFileIndex {
pub file_id: u32,
pub hashes: Vec<[u8; 32]>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ResumeEntry {
pub file_id: u32,
pub have: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Control {
Manifest(Vec<FileEntry>),
ResumeState(Vec<ResumeEntry>),
LocalIndex(Vec<LocalFileIndex>),
Start {
streams: u32,
},
FileComplete {
file_id: u32,
hash: Option<[u8; 32]>,
},
AllComplete,
Abort {
reason: String,
},
}
impl Control {
fn kind(&self) -> ControlKind {
match self {
Control::LocalIndex(_) => ControlKind::LocalIndex,
Control::Manifest(_) => ControlKind::Manifest,
Control::ResumeState(_) => ControlKind::ResumeState,
Control::Start { .. } => ControlKind::Start,
Control::FileComplete { .. } => ControlKind::FileComplete,
Control::AllComplete => ControlKind::AllComplete,
Control::Abort { .. } => ControlKind::Abort,
}
}
fn encode_body(&self, out: &mut Vec<u8>) {
match self {
Control::Manifest(entries) => {
out.extend_from_slice(&(entries.len() as u32).to_le_bytes());
for e in entries {
e.encode(out);
}
}
Control::ResumeState(entries) => {
out.extend_from_slice(&(entries.len() as u32).to_le_bytes());
for e in entries {
out.extend_from_slice(&e.file_id.to_le_bytes());
out.extend_from_slice(&(e.have.len() as u32).to_le_bytes());
out.extend_from_slice(&e.have);
}
}
Control::LocalIndex(entries) => {
out.extend_from_slice(&(entries.len() as u32).to_le_bytes());
for e in entries {
out.extend_from_slice(&e.file_id.to_le_bytes());
out.extend_from_slice(&(e.hashes.len() as u32).to_le_bytes());
for h in &e.hashes {
out.extend_from_slice(h);
}
}
}
Control::AllComplete => {}
Control::Start { streams } => {
out.extend_from_slice(&streams.to_le_bytes());
}
Control::FileComplete { file_id, hash } => {
out.extend_from_slice(&file_id.to_le_bytes());
match hash {
Some(h) => {
out.push(1);
out.extend_from_slice(h);
}
None => out.push(0),
}
}
Control::Abort { reason } => {
let r = reason.as_bytes();
let n = r.len().min(1024);
out.extend_from_slice(&(n as u16).to_le_bytes());
out.extend_from_slice(&r[..n]);
}
}
}
fn decode_body(kind: ControlKind, body: &[u8], max_entries: usize) -> Result<Self> {
let mut cur = Cursor::new(body);
Ok(match kind {
ControlKind::Manifest => {
let n = cur.u32()? as usize;
if n > max_entries {
return Err(Error::protocol(format!(
"manifest declares {n} entries, limit is {max_entries}"
)));
}
let mut v = Vec::with_capacity(n.min(4096));
for _ in 0..n {
v.push(FileEntry::decode(&mut cur)?);
}
Control::Manifest(v)
}
ControlKind::ResumeState => {
let n = cur.u32()? as usize;
if n > max_entries {
return Err(Error::protocol("resume state entry count exceeds limit"));
}
let mut v = Vec::with_capacity(n.min(4096));
for _ in 0..n {
let file_id = cur.u32()?;
let len = cur.u32()? as usize;
let have = cur.take(len)?.to_vec();
v.push(ResumeEntry { file_id, have });
}
Control::ResumeState(v)
}
ControlKind::LocalIndex => {
let n = cur.u32()? as usize;
if n > max_entries {
return Err(Error::protocol("local index entry count exceeds limit"));
}
let mut v = Vec::with_capacity(n.min(4096));
for _ in 0..n {
let file_id = cur.u32()?;
let count = cur.u32()? as usize;
if count > (1 << 26) {
return Err(Error::protocol("local index declares too many chunks"));
}
let mut hashes = Vec::with_capacity(count.min(4096));
for _ in 0..count {
let mut h = [0u8; 32];
h.copy_from_slice(cur.take(32)?);
hashes.push(h);
}
v.push(LocalFileIndex { file_id, hashes });
}
Control::LocalIndex(v)
}
ControlKind::Start => Control::Start {
streams: cur.u32()?,
},
ControlKind::AllComplete => Control::AllComplete,
ControlKind::FileComplete => {
let file_id = cur.u32()?;
let hash = if cur.u8()? == 1 {
let mut h = [0u8; 32];
h.copy_from_slice(cur.take(32)?);
Some(h)
} else {
None
};
Control::FileComplete { file_id, hash }
}
ControlKind::Abort => {
let n = cur.u16()? as usize;
let s = cur.take(n)?;
Control::Abort {
reason: String::from_utf8_lossy(s).into_owned(),
}
}
})
}
}
pub async fn write_control<W: AsyncWrite + Unpin>(w: &mut W, msg: &Control) -> Result<()> {
let mut body = Vec::with_capacity(64);
msg.encode_body(&mut body);
let mut framed = Vec::with_capacity(body.len() + 5);
framed.push(msg.kind() as u8);
framed.extend_from_slice(&(body.len() as u32).to_le_bytes());
framed.extend_from_slice(&body);
w.write_all(&framed).await?;
w.flush().await?;
Ok(())
}
pub async fn read_control<R: AsyncRead + Unpin>(
r: &mut R,
max_body: usize,
max_entries: usize,
) -> Result<Control> {
let mut hdr = [0u8; 5];
r.read_exact(&mut hdr).await.map_err(map_eof)?;
let kind = ControlKind::from_u8(hdr[0])?;
let len = u32::from_le_bytes(hdr[1..5].try_into().unwrap()) as usize;
if len > max_body {
return Err(Error::FrameTooLarge {
got: len,
limit: max_body,
});
}
let mut body = vec![0u8; len];
r.read_exact(&mut body).await.map_err(map_eof)?;
Control::decode_body(kind, &body, max_entries)
}
fn map_eof(e: std::io::Error) -> Error {
if e.kind() == std::io::ErrorKind::UnexpectedEof {
Error::Closed("stream ended mid-message".into())
} else {
Error::Io(e)
}
}
struct Cursor<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> Cursor<'a> {
fn new(buf: &'a [u8]) -> Self {
Self { buf, pos: 0 }
}
fn take(&mut self, n: usize) -> Result<&'a [u8]> {
let end = self
.pos
.checked_add(n)
.ok_or_else(|| Error::protocol("length overflow"))?;
if end > self.buf.len() {
return Err(Error::protocol("message truncated"));
}
let s = &self.buf[self.pos..end];
self.pos = end;
Ok(s)
}
fn u8(&mut self) -> Result<u8> {
Ok(self.take(1)?[0])
}
fn u16(&mut self) -> Result<u16> {
Ok(u16::from_le_bytes(self.take(2)?.try_into().unwrap()))
}
fn u32(&mut self) -> Result<u32> {
Ok(u32::from_le_bytes(self.take(4)?.try_into().unwrap()))
}
fn u64(&mut self) -> Result<u64> {
Ok(u64::from_le_bytes(self.take(8)?.try_into().unwrap()))
}
fn i64(&mut self) -> Result<i64> {
Ok(i64::from_le_bytes(self.take(8)?.try_into().unwrap()))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_entry(id: u32) -> FileEntry {
FileEntry {
file_id: id,
path: format!("music/album {id}/track.flac"),
size: 987_654_321_000,
chunk_size: 1 << 20,
mode: 0o644,
mtime: 1_700_000_000,
kind: EntryKind::File,
hash: Some([id as u8; 32]),
incompressible: true,
}
}
#[test]
fn frame_header_roundtrip() {
let h = FrameHeader {
flags: flags::LAST_CHUNK | flags::SEALED,
algorithm: Algorithm::Zstd,
file_id: 0xDEAD_BEEF,
chunk_index: 9_876_543_210,
epoch: 3,
raw_len: 1 << 20,
payload_len: 700_000,
};
let mut buf = [0u8; FRAME_HEADER_LEN];
h.encode(&mut buf);
let d = FrameHeader::decode(&buf, 64 << 20).unwrap();
assert_eq!(h, d);
assert!(d.sealed() && d.last_chunk());
assert_eq!(d.wire_payload_len(), 700_000 + 16);
}
#[test]
fn frame_header_rejects_oversized_lengths() {
let h = FrameHeader {
flags: 0,
algorithm: Algorithm::None,
file_id: 1,
chunk_index: 0,
epoch: 0,
raw_len: 16,
payload_len: u32::MAX,
};
let mut buf = [0u8; FRAME_HEADER_LEN];
h.encode(&mut buf);
assert!(matches!(
FrameHeader::decode(&buf, 1 << 20),
Err(Error::FrameTooLarge { .. })
));
}
#[test]
fn frame_header_rejects_bad_magic_and_version() {
let mut buf = [0u8; FRAME_HEADER_LEN];
FrameHeader {
flags: 0,
algorithm: Algorithm::None,
file_id: 1,
chunk_index: 0,
epoch: 0,
raw_len: 1,
payload_len: 1,
}
.encode(&mut buf);
let mut bad = buf;
bad[0] = 0x00;
assert!(FrameHeader::decode(&bad, 1 << 20).is_err());
let mut bad = buf;
bad[1] = 99;
assert!(matches!(
FrameHeader::decode(&bad, 1 << 20),
Err(Error::Version { .. })
));
}
#[tokio::test]
async fn control_roundtrip_all_variants() {
let msgs = vec![
Control::Manifest((0..64).map(sample_entry).collect()),
Control::ResumeState(vec![ResumeEntry {
file_id: 3,
have: vec![0xFF, 0x0F, 0x00],
}]),
Control::LocalIndex(vec![LocalFileIndex {
file_id: 4,
hashes: vec![[1u8; 32], [2u8; 32], [3u8; 32]],
}]),
Control::Start { streams: 16 },
Control::FileComplete {
file_id: 12,
hash: Some([7u8; 32]),
},
Control::FileComplete {
file_id: 13,
hash: None,
},
Control::AllComplete,
Control::Abort {
reason: "disk full".into(),
},
];
for m in msgs {
let mut buf = Vec::new();
write_control(&mut buf, &m).await.unwrap();
let mut slice = &buf[..];
let got = read_control(&mut slice, 16 << 20, 1 << 20).await.unwrap();
assert_eq!(m, got);
}
}
#[tokio::test]
async fn control_rejects_oversized_body() {
let m = Control::Manifest((0..1000).map(sample_entry).collect());
let mut buf = Vec::new();
write_control(&mut buf, &m).await.unwrap();
let mut slice = &buf[..];
assert!(matches!(
read_control(&mut slice, 128, 1 << 20).await,
Err(Error::FrameTooLarge { .. })
));
}
#[tokio::test]
async fn control_rejects_absurd_entry_count() {
let mut body = Vec::new();
body.extend_from_slice(&(5_000_000u32).to_le_bytes());
let mut framed = vec![ControlKind::Manifest as u8];
framed.extend_from_slice(&(body.len() as u32).to_le_bytes());
framed.extend_from_slice(&body);
let mut slice = &framed[..];
assert!(read_control(&mut slice, 16 << 20, 1000).await.is_err());
}
#[tokio::test]
async fn truncated_manifest_is_an_error_not_a_panic() {
let m = Control::Manifest((0..8).map(sample_entry).collect());
let mut buf = Vec::new();
write_control(&mut buf, &m).await.unwrap();
let cut = buf.len() - 40;
let body_len = u32::from_le_bytes(buf[1..5].try_into().unwrap()) as usize;
let truncated_body = &buf[5..cut];
let r = Control::decode_body(ControlKind::Manifest, truncated_body, 1 << 20);
assert!(r.is_err());
assert!(body_len > truncated_body.len());
}
#[test]
fn chunk_count_is_exact_at_boundaries() {
let mut e = sample_entry(1);
e.chunk_size = 1024;
e.size = 0;
assert_eq!(e.chunk_count(), 0);
e.size = 1;
assert_eq!(e.chunk_count(), 1);
e.size = 1024;
assert_eq!(e.chunk_count(), 1);
e.size = 1025;
assert_eq!(e.chunk_count(), 2);
e.chunk_size = 1 << 20;
e.size = 100 * (1 << 30);
assert_eq!(e.chunk_count(), 102_400);
}
}