use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use tracing::trace;
pub const DEFAULT_MAX_DATAGRAM: usize = 1200;
pub const PROTOCOL_VERSION: u32 = 1;
pub const FRAGMENT_HEADER_OVERHEAD: usize = 10;
pub const MAX_FRAGMENT_INDEX: u16 = 0x7fff;
#[derive(Debug, thiserror::Error)]
pub enum WireError {
#[error("postcard (de)serialization error: {0}")]
Postcard(#[from] postcard::Error),
#[error("MTU {mtu} too small for fragment header (need > {min})")]
MtuTooSmall { mtu: usize, min: usize },
#[error("protocol version mismatch: peer sent {peer}, we speak {ours}")]
VersionMismatch { peer: u32, ours: u32 },
#[error("fragment too short: {len} bytes (need >= {min})")]
ShortFragment { len: usize, min: usize },
#[error("instruction needs {count} fragments, exceeds the {max}-fragment limit")]
TooManyFragments { count: usize, max: usize },
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Instruction {
pub protocol_version: u32,
pub old_num: u64,
pub new_num: u64,
pub ack_num: u64,
pub throwaway_num: u64,
pub diff: Vec<u8>,
}
impl Instruction {
pub const fn ack_only(state_num: u64, ack_num: u64, throwaway_num: u64) -> Self {
Self {
protocol_version: PROTOCOL_VERSION,
old_num: state_num,
new_num: state_num,
ack_num,
throwaway_num,
diff: Vec::new(),
}
}
pub fn encode(&self) -> Result<Vec<u8>, WireError> {
Ok(postcard::to_allocvec(self)?)
}
pub fn decode(bytes: &[u8]) -> Result<Self, WireError> {
let instr: Self = postcard::from_bytes(bytes)?;
if instr.protocol_version != PROTOCOL_VERSION {
return Err(WireError::VersionMismatch {
peer: instr.protocol_version,
ours: PROTOCOL_VERSION,
});
}
Ok(instr)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Fragment {
pub id: u64,
pub index: u16,
pub final_: bool,
pub payload: Vec<u8>,
}
impl Fragment {
pub fn encode(&self) -> Result<Vec<u8>, WireError> {
let combined: u16 = (u16::from(self.final_) << 15) | (self.index & MAX_FRAGMENT_INDEX);
let mut out = Vec::with_capacity(FRAGMENT_HEADER_OVERHEAD + self.payload.len());
out.extend_from_slice(&self.id.to_be_bytes());
out.extend_from_slice(&combined.to_be_bytes());
out.extend_from_slice(&self.payload);
Ok(out)
}
pub fn decode(bytes: &[u8]) -> Result<Self, WireError> {
let short = || WireError::ShortFragment {
len: bytes.len(),
min: FRAGMENT_HEADER_OVERHEAD,
};
let (id_bytes, rest) = bytes.split_first_chunk::<8>().ok_or_else(short)?;
let (combined_bytes, payload) = rest.split_first_chunk::<2>().ok_or_else(short)?;
let id = u64::from_be_bytes(*id_bytes);
let combined = u16::from_be_bytes(*combined_bytes);
Ok(Self {
id,
index: combined & MAX_FRAGMENT_INDEX,
final_: combined & 0x8000 != 0,
payload: payload.to_vec(),
})
}
}
#[derive(Debug, Default)]
pub struct Fragmenter {
next_id: u64,
last_serialized: Option<Vec<u8>>,
last_mtu: usize,
}
impl Fragmenter {
pub fn new() -> Self {
Self::default()
}
pub const fn current_id(&self) -> u64 {
self.next_id
}
pub fn fragment(
&mut self,
instr: &Instruction,
mtu: usize,
) -> Result<Vec<Fragment>, WireError> {
if mtu <= FRAGMENT_HEADER_OVERHEAD {
return Err(WireError::MtuTooSmall {
mtu,
min: FRAGMENT_HEADER_OVERHEAD,
});
}
let serialized = instr.encode()?;
let changed = match &self.last_serialized {
Some(prev) => prev != &serialized || self.last_mtu != mtu,
None => true,
};
if changed {
self.next_id = self.next_id.wrapping_add(1);
self.last_serialized = Some(serialized.clone());
self.last_mtu = mtu;
}
let id = self.next_id;
let chunk = mtu - FRAGMENT_HEADER_OVERHEAD;
let mut fragments = Vec::new();
if serialized.is_empty() {
trace!(
id,
changed,
mtu,
"fragmented empty instruction into 1 fragment"
);
fragments.push(Fragment {
id,
index: 0,
final_: true,
payload: Vec::new(),
});
return Ok(fragments);
}
let total = serialized.len().div_ceil(chunk);
if total > MAX_FRAGMENT_INDEX as usize + 1 {
return Err(WireError::TooManyFragments {
count: total,
max: MAX_FRAGMENT_INDEX as usize + 1,
});
}
trace!(
id,
fragments = total,
bytes = serialized.len(),
mtu,
changed,
"fragmented instruction"
);
for (i, piece) in serialized.chunks(chunk).enumerate() {
fragments.push(Fragment {
id,
index: i as u16,
final_: i + 1 == total,
payload: piece.to_vec(),
});
}
Ok(fragments)
}
}
#[derive(Debug, Default)]
pub struct FragmentAssembly {
current_id: Option<u64>,
parts: BTreeMap<u16, Vec<u8>>,
final_index: Option<u16>,
}
impl FragmentAssembly {
pub fn new() -> Self {
Self::default()
}
pub fn add(&mut self, frag: Fragment) -> Result<Option<Instruction>, WireError> {
match self.current_id {
Some(cur) if frag.id < cur => {
trace!(
stale_id = frag.id,
current = cur,
"dropping superseded fragment"
);
return Ok(None); }
Some(cur) if frag.id == cur => {} _ => {
if let Some(prev) = self.current_id {
trace!(
prev_id = prev,
new_id = frag.id,
"newer instruction supersedes partial"
);
}
self.current_id = Some(frag.id);
self.parts.clear();
self.final_index = None;
}
}
if frag.final_ {
self.final_index = Some(frag.index);
}
self.parts.insert(frag.index, frag.payload);
let Some(final_idx) = self.final_index else {
return Ok(None);
};
let needed = final_idx as usize + 1;
if self.parts.len() != needed || self.parts.keys().next_back() != Some(&final_idx) {
return Ok(None);
}
let mut buf = Vec::new();
for i in 0..=final_idx {
match self.parts.get(&i) {
Some(part) => buf.extend_from_slice(part),
None => return Ok(None), }
}
self.parts.clear();
self.final_index = None;
trace!(
id = self.current_id,
fragments = needed,
"reassembled complete instruction"
);
let instr = Instruction::decode(&buf)?;
Ok(Some(instr))
}
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
fn sample_instruction(diff_len: usize) -> Instruction {
Instruction {
protocol_version: PROTOCOL_VERSION,
old_num: 7,
new_num: 9,
ack_num: 4,
throwaway_num: 2,
diff: (0..diff_len).map(|i| (i % 251) as u8).collect(),
}
}
#[test]
fn instruction_roundtrip() {
let i = sample_instruction(300);
let bytes = i.encode().unwrap();
assert_eq!(Instruction::decode(&bytes).unwrap(), i);
}
#[test]
fn decode_rejects_protocol_version_mismatch() {
let mut i = sample_instruction(10);
i.protocol_version = PROTOCOL_VERSION + 1;
let bytes = i.encode().unwrap();
match Instruction::decode(&bytes) {
Err(WireError::VersionMismatch { peer, ours }) => {
assert_eq!(peer, PROTOCOL_VERSION + 1);
assert_eq!(ours, PROTOCOL_VERSION);
}
other => panic!("expected VersionMismatch, got {other:?}"),
}
let ok = sample_instruction(10);
assert_eq!(Instruction::decode(&ok.encode().unwrap()).unwrap(), ok);
}
#[test]
fn small_instruction_single_fragment() {
let mut f = Fragmenter::new();
let frags = f.fragment(&sample_instruction(10), 1200).unwrap();
assert_eq!(frags.len(), 1);
assert!(frags[0].final_);
}
#[test]
fn empty_diff_yields_one_fragment() {
let mut f = Fragmenter::new();
let instr = Instruction::ack_only(5, 3, 1);
let frags = f.fragment(&instr, 1200).unwrap();
assert_eq!(frags.len(), 1);
assert!(frags[0].final_);
assert!(instr.diff.is_empty());
let mut asm = FragmentAssembly::new();
assert_eq!(asm.add(frags[0].clone()).unwrap().unwrap(), instr);
}
#[test]
fn large_instruction_fragments_and_reassembles() {
let mut f = Fragmenter::new();
let instr = sample_instruction(10_000);
let frags = f.fragment(&instr, 200).unwrap();
assert!(frags.len() > 1);
for w in frags.windows(2) {
assert_eq!(w[1].index, w[0].index + 1);
}
assert!(frags.last().unwrap().final_);
for fr in &frags {
assert!(fr.encode().unwrap().len() <= 200, "fragment exceeds MTU");
}
let mut asm = FragmentAssembly::new();
let mut got = None;
for fr in frags {
if let Some(i) = asm.add(fr).unwrap() {
got = Some(i);
}
}
assert_eq!(got.unwrap(), instr);
}
#[test]
fn fragment_header_is_exact_and_packs_to_mtu() {
assert_eq!(FRAGMENT_HEADER_OVERHEAD, 10);
let f = Fragment {
id: 0xdead_beef_cafe_babe,
index: 5,
final_: true,
payload: vec![7u8; 100],
};
let bytes = f.encode().unwrap();
assert_eq!(bytes.len(), FRAGMENT_HEADER_OVERHEAD + 100);
assert_eq!(Fragment::decode(&bytes).unwrap(), f);
let mut fr = Fragmenter::new();
let frags = fr.fragment(&sample_instruction(5000), 200).unwrap();
assert!(frags.len() > 1);
for f in &frags[..frags.len() - 1] {
assert_eq!(
f.encode().unwrap().len(),
200,
"non-final fragments pack to exactly the MTU"
);
}
}
#[test]
fn decode_rejects_short_fragment() {
assert!(matches!(
Fragment::decode(&[0u8; 5]),
Err(WireError::ShortFragment { .. })
));
}
#[test]
fn identical_retransmit_reuses_id() {
let mut f = Fragmenter::new();
let instr = sample_instruction(50);
let a = f.fragment(&instr, 1200).unwrap();
let b = f.fragment(&instr, 1200).unwrap();
assert_eq!(a[0].id, b[0].id, "identical content must reuse fragment id");
let mut instr2 = instr.clone();
instr2.new_num += 1;
let c = f.fragment(&instr2, 1200).unwrap();
assert!(c[0].id > a[0].id, "changed content must bump fragment id");
}
#[test]
fn newer_id_supersedes_partial() {
let mut f = Fragmenter::new();
let old = sample_instruction(5_000);
let old_frags = f.fragment(&old, 300).unwrap();
let mut newer = sample_instruction(5_000);
newer.new_num = 999;
let new_frags = f.fragment(&newer, 300).unwrap();
let mut asm = FragmentAssembly::new();
assert!(asm.add(old_frags[0].clone()).unwrap().is_none());
let mut got = None;
for fr in new_frags {
if let Some(i) = asm.add(fr).unwrap() {
got = Some(i);
}
}
assert_eq!(got.unwrap(), newer);
}
#[test]
fn high_index_partial_does_not_complete_and_is_superseded() {
let mut asm = FragmentAssembly::new();
let high = Fragment {
id: 5,
index: MAX_FRAGMENT_INDEX,
final_: false,
payload: vec![1u8; 8],
};
assert!(
asm.add(high).unwrap().is_none(),
"a lone non-final fragment can never complete"
);
let instr = sample_instruction(20);
let frags = Fragmenter::new().fragment(&instr, 1200).unwrap();
let mut got = None;
for mut fr in frags {
fr.id = 6; if let Some(i) = asm.add(fr).unwrap() {
got = Some(i);
}
}
assert_eq!(
got.unwrap(),
instr,
"the newer instruction supersedes the stale high-index partial"
);
}
proptest! {
#[test]
fn fragment_reassemble_roundtrip(diff_len in 0usize..20_000, mtu in 30usize..1500) {
let mut f = Fragmenter::new();
let instr = Instruction {
protocol_version: PROTOCOL_VERSION,
old_num: 1, new_num: 2, ack_num: 0, throwaway_num: 0,
diff: (0..diff_len).map(|i| (i % 256) as u8).collect(),
};
let frags = f.fragment(&instr, mtu).unwrap();
for fr in &frags {
prop_assert!(fr.encode().unwrap().len() <= mtu);
}
let mut asm = FragmentAssembly::new();
let mut got = None;
for fr in frags {
if let Some(i) = asm.add(fr).unwrap() { got = Some(i); }
}
prop_assert_eq!(got.unwrap(), instr);
}
}
}