From 5d13090d8f2e5f7d10bc0493a7fca5b1202a258a Mon Sep 17 00:00:00 2001 From: Johnathan Corgan Date: Wed, 8 Jul 2026 16:37:17 +0000 Subject: [PATCH] proto: add a bounds-checked byte reader/writer and adopt it across the wire codecs Introduce a shared proto/codec module with a cursor Reader (short reads fail with MessageTooShort { expected: position + needed, got: total }) and an append Writer, and adopt them across the seven subsystem wire codecs, replacing the repetitive manual slicing, try_into, and from_le_bytes extraction. Each existing length check maps to the reader (an up-front minimum becomes require at position zero, so the per-field expected values, including the tree-announce expected 99, are reproduced exactly); the bloom exact-length check stays explicit since it also rejects over-long payloads. Encoded bytes and decode decisions are unchanged. --- src/proto/bloom/wire.rs | 47 +++---- src/proto/codec.rs | 109 +++++++++++++++ src/proto/discovery/wire.rs | 81 +++--------- src/proto/fmp/wire.rs | 11 +- src/proto/fsp/wire.rs | 119 +++++------------ src/proto/mmp/wire.rs | 257 +++++++++++++++++------------------- src/proto/mod.rs | 1 + src/proto/routing/wire.rs | 88 ++++++------ src/proto/stp/wire.rs | 98 ++++---------- 9 files changed, 365 insertions(+), 446 deletions(-) create mode 100644 src/proto/codec.rs diff --git a/src/proto/bloom/wire.rs b/src/proto/bloom/wire.rs index e180982..0b3c8a3 100644 --- a/src/proto/bloom/wire.rs +++ b/src/proto/bloom/wire.rs @@ -2,6 +2,7 @@ use super::BloomFilter; use crate::proto::Error; +use crate::proto::codec::{Reader, Writer}; use crate::proto::link::LinkMessageType; /// Bloom filter announcement for reachability propagation. @@ -86,50 +87,37 @@ impl FilterAnnounce { let filter_bytes = self.filter.as_bytes(); let size = 1 + Self::MIN_PAYLOAD_SIZE + filter_bytes.len(); - let mut buf = Vec::with_capacity(size); + let mut w = Writer::with_capacity(size); // msg_type - buf.push(LinkMessageType::FilterAnnounce.to_byte()); + w.write_u8(LinkMessageType::FilterAnnounce.to_byte()); // sequence (8 LE) - buf.extend_from_slice(&self.sequence.to_le_bytes()); + w.write_u64_le(self.sequence); // hash_count - buf.push(self.hash_count); + w.write_u8(self.hash_count); // size_class - buf.push(self.size_class); + w.write_u8(self.size_class); // filter_bits - buf.extend_from_slice(filter_bytes); + w.write_bytes(filter_bytes); - Ok(buf) + Ok(w.into_vec()) } /// Decode from link-layer payload (after msg_type byte stripped by dispatcher). /// /// The payload starts with the sequence field. pub fn decode(payload: &[u8]) -> Result { - if payload.len() < Self::MIN_PAYLOAD_SIZE { - return Err(Error::MessageTooShort { - expected: Self::MIN_PAYLOAD_SIZE, - got: payload.len(), - }); - } - - let mut pos = 0; + let mut reader = Reader::new(payload); + reader.require(Self::MIN_PAYLOAD_SIZE)?; // sequence (8 LE) - let sequence = u64::from_le_bytes( - payload[pos..pos + 8] - .try_into() - .map_err(|_| Error::Malformed("bad sequence"))?, - ); - pos += 8; + let sequence = reader.read_u64_le()?; // hash_count - let hash_count = payload[pos]; - pos += 1; + let hash_count = reader.read_u8()?; // size_class - let size_class = payload[pos]; - pos += 1; + let size_class = reader.read_u8()?; // Validate size_class range if size_class > Self::MAX_SIZE_CLASS { @@ -147,9 +135,11 @@ impl FilterAnnounce { }); } - // Expected filter size from size_class + // Expected filter size from size_class. The remaining length must match + // exactly (an over-long payload is rejected too), so this stays an + // explicit `!=` check rather than a `Reader::require` lower-bound gate. let expected_filter_bytes = 512usize << size_class; - let remaining = payload.len() - pos; + let remaining = reader.remaining(); if remaining != expected_filter_bytes { return Err(Error::MessageTooShort { expected: Self::MIN_PAYLOAD_SIZE + expected_filter_bytes, @@ -158,8 +148,7 @@ impl FilterAnnounce { } // Construct BloomFilter from bytes - let filter = - BloomFilter::from_slice(&payload[pos..], hash_count).map_err(Error::BadBloom)?; + let filter = BloomFilter::from_slice(reader.rest(), hash_count).map_err(Error::BadBloom)?; let announce = Self { filter, diff --git a/src/proto/codec.rs b/src/proto/codec.rs new file mode 100644 index 0000000..0ca0a97 --- /dev/null +++ b/src/proto/codec.rs @@ -0,0 +1,109 @@ +//! Bounds-checked byte reader/writer shared across the proto wire codecs. +//! +//! `Reader` fails a short read with `Error::MessageTooShort { expected, got }` +//! where `expected` is the cumulative byte offset it needed (`position + n`) and +//! `got` is the total buffer length — reproducing the codecs' existing per-field +//! and up-front length-check values exactly. +use crate::proto::Error; + +pub(crate) struct Reader<'a> { + buf: &'a [u8], + pos: usize, +} + +impl<'a> Reader<'a> { + pub(crate) fn new(buf: &'a [u8]) -> Self { + Self { buf, pos: 0 } + } + #[allow(dead_code)] + pub(crate) fn position(&self) -> usize { + self.pos + } + pub(crate) fn remaining(&self) -> usize { + self.buf.len() - self.pos + } + pub(crate) fn rest(&self) -> &'a [u8] { + &self.buf[self.pos..] + } + /// Ensure at least `n` more bytes are available; else MessageTooShort. + pub(crate) fn require(&self, n: usize) -> Result<(), Error> { + if self.pos + n > self.buf.len() { + return Err(Error::MessageTooShort { + expected: self.pos + n, + got: self.buf.len(), + }); + } + Ok(()) + } + /// Advance the cursor by `n` (caller has already validated bounds, e.g. via a + /// sub-decoder that returned a consumed count). Debug-panics if out of range. + pub(crate) fn advance(&mut self, n: usize) { + self.pos += n; + debug_assert!(self.pos <= self.buf.len()); + } + pub(crate) fn read_u8(&mut self) -> Result { + self.require(1)?; + let v = self.buf[self.pos]; + self.pos += 1; + Ok(v) + } + pub(crate) fn read_array(&mut self) -> Result<[u8; N], Error> { + self.require(N)?; + let mut a = [0u8; N]; + a.copy_from_slice(&self.buf[self.pos..self.pos + N]); + self.pos += N; + Ok(a) + } + pub(crate) fn read_u16_le(&mut self) -> Result { + Ok(u16::from_le_bytes(self.read_array::<2>()?)) + } + pub(crate) fn read_u32_le(&mut self) -> Result { + Ok(u32::from_le_bytes(self.read_array::<4>()?)) + } + pub(crate) fn read_u64_le(&mut self) -> Result { + Ok(u64::from_le_bytes(self.read_array::<8>()?)) + } + pub(crate) fn read_bytes(&mut self, n: usize) -> Result<&'a [u8], Error> { + self.require(n)?; + let s = &self.buf[self.pos..self.pos + n]; + self.pos += n; + Ok(s) + } +} + +pub(crate) struct Writer { + buf: alloc::vec::Vec, +} +impl Writer { + pub(crate) fn new() -> Self { + Self { + buf: alloc::vec::Vec::new(), + } + } + pub(crate) fn with_capacity(n: usize) -> Self { + Self { + buf: alloc::vec::Vec::with_capacity(n), + } + } + pub(crate) fn write_u8(&mut self, v: u8) { + self.buf.push(v); + } + pub(crate) fn write_u16_le(&mut self, v: u16) { + self.buf.extend_from_slice(&v.to_le_bytes()); + } + pub(crate) fn write_u32_le(&mut self, v: u32) { + self.buf.extend_from_slice(&v.to_le_bytes()); + } + pub(crate) fn write_u64_le(&mut self, v: u64) { + self.buf.extend_from_slice(&v.to_le_bytes()); + } + pub(crate) fn write_bytes(&mut self, b: &[u8]) { + self.buf.extend_from_slice(b); + } + pub(crate) fn len(&self) -> usize { + self.buf.len() + } + pub(crate) fn into_vec(self) -> alloc::vec::Vec { + self.buf + } +} diff --git a/src/proto/discovery/wire.rs b/src/proto/discovery/wire.rs index c717af4..c31c539 100644 --- a/src/proto/discovery/wire.rs +++ b/src/proto/discovery/wire.rs @@ -2,6 +2,7 @@ use crate::NodeAddr; use crate::proto::Error; +use crate::proto::codec::Reader; use crate::proto::stp::TreeCoordinate; use crate::proto::stp::{decode_coords, encode_coords}; use secp256k1::schnorr::Signature; @@ -86,43 +87,20 @@ impl LookupRequest { pub fn decode(payload: &[u8]) -> Result { // Minimum: request_id(8) + target(16) + origin(16) + ttl(1) + min_mtu(2) // + coords_count(2) = 45 bytes - if payload.len() < 45 { - return Err(Error::MessageTooShort { - expected: 45, - got: payload.len(), - }); - } + let mut reader = Reader::new(payload); + reader.require(45)?; - let mut pos = 0; + let request_id = reader.read_u64_le()?; - let request_id = u64::from_le_bytes( - payload[pos..pos + 8] - .try_into() - .map_err(|_| Error::Malformed("bad request_id"))?, - ); - pos += 8; + let target = NodeAddr::from_bytes(reader.read_array::<16>()?); - let mut target_bytes = [0u8; 16]; - target_bytes.copy_from_slice(&payload[pos..pos + 16]); - let target = NodeAddr::from_bytes(target_bytes); - pos += 16; + let origin = NodeAddr::from_bytes(reader.read_array::<16>()?); - let mut origin_bytes = [0u8; 16]; - origin_bytes.copy_from_slice(&payload[pos..pos + 16]); - let origin = NodeAddr::from_bytes(origin_bytes); - pos += 16; + let ttl = reader.read_u8()?; - let ttl = payload[pos]; - pos += 1; + let min_mtu = reader.read_u16_le()?; - let min_mtu = u16::from_le_bytes( - payload[pos..pos + 2] - .try_into() - .map_err(|_| Error::Malformed("bad min_mtu"))?, - ); - pos += 2; - - let (origin_coords, _consumed) = decode_coords(&payload[pos..])?; + let (origin_coords, _consumed) = decode_coords(reader.rest())?; Ok(Self { request_id, @@ -211,44 +189,19 @@ impl LookupResponse { /// Decode from wire format (after msg_type byte has been consumed). pub fn decode(payload: &[u8]) -> Result { // Minimum: request_id(8) + target(16) + path_mtu(2) + coords_count(2) + proof(64) = 92 - if payload.len() < 92 { - return Err(Error::MessageTooShort { - expected: 92, - got: payload.len(), - }); - } + let mut reader = Reader::new(payload); + reader.require(92)?; - let mut pos = 0; + let request_id = reader.read_u64_le()?; - let request_id = u64::from_le_bytes( - payload[pos..pos + 8] - .try_into() - .map_err(|_| Error::Malformed("bad request_id"))?, - ); - pos += 8; + let target = NodeAddr::from_bytes(reader.read_array::<16>()?); - let mut target_bytes = [0u8; 16]; - target_bytes.copy_from_slice(&payload[pos..pos + 16]); - let target = NodeAddr::from_bytes(target_bytes); - pos += 16; + let path_mtu = reader.read_u16_le()?; - let path_mtu = u16::from_le_bytes( - payload[pos..pos + 2] - .try_into() - .map_err(|_| Error::Malformed("bad path_mtu"))?, - ); - pos += 2; + let (target_coords, consumed) = decode_coords(reader.rest())?; + reader.advance(consumed); - let (target_coords, consumed) = decode_coords(&payload[pos..])?; - pos += consumed; - - if payload.len() < pos + 64 { - return Err(Error::MessageTooShort { - expected: pos + 64, - got: payload.len(), - }); - } - let proof = Signature::from_slice(&payload[pos..pos + 64]) + let proof = Signature::from_slice(reader.read_bytes(64)?) .map_err(|_| Error::Malformed("bad proof signature"))?; Ok(Self { diff --git a/src/proto/fmp/wire.rs b/src/proto/fmp/wire.rs index 0da6fd3..b9d8b6c 100644 --- a/src/proto/fmp/wire.rs +++ b/src/proto/fmp/wire.rs @@ -7,6 +7,7 @@ //! in `crate::proto::link`. use crate::proto::Error; +use crate::proto::codec::Reader; use crate::proto::link::LinkMessageType; use std::fmt; @@ -153,13 +154,9 @@ impl Disconnect { /// Decode from link-layer payload (after msg_type byte has been consumed). pub fn decode(payload: &[u8]) -> Result { - if payload.is_empty() { - return Err(Error::MessageTooShort { - expected: 1, - got: 0, - }); - } - let reason = DisconnectReason::from_byte(payload[0]).unwrap_or(DisconnectReason::Other); + let mut reader = Reader::new(payload); + let reason = + DisconnectReason::from_byte(reader.read_u8()?).unwrap_or(DisconnectReason::Other); Ok(Self { reason }) } } diff --git a/src/proto/fsp/wire.rs b/src/proto/fsp/wire.rs index b2ec872..ce0e933 100644 --- a/src/proto/fsp/wire.rs +++ b/src/proto/fsp/wire.rs @@ -33,6 +33,7 @@ //! | 0x3 | - | Handshake msg3 | SessionMsg3 (Noise XK msg3) | use crate::proto::Error; +use crate::proto::codec::{Reader, Writer}; use crate::proto::stp::{TreeCoordinate, decode_coords, decode_optional_coords, encode_coords}; use std::fmt; @@ -647,37 +648,18 @@ impl SessionSetup { /// Decode from wire format (after 4-byte FSP prefix has been consumed). pub fn decode(payload: &[u8]) -> Result { - if payload.is_empty() { - return Err(Error::MessageTooShort { - expected: 1, - got: 0, - }); - } - let flags = SessionFlags::from_byte(payload[0]); - let mut offset = 1; + let mut reader = Reader::new(payload); + let flags = SessionFlags::from_byte(reader.read_u8()?); - let (src_coords, consumed) = decode_coords(&payload[offset..])?; - offset += consumed; + let (src_coords, consumed) = decode_coords(reader.rest())?; + reader.advance(consumed); - let (dest_coords, consumed) = decode_coords(&payload[offset..])?; - offset += consumed; + let (dest_coords, consumed) = decode_coords(reader.rest())?; + reader.advance(consumed); - if payload.len() < offset + 2 { - return Err(Error::MessageTooShort { - expected: offset + 2, - got: payload.len(), - }); - } - let hs_len = u16::from_le_bytes([payload[offset], payload[offset + 1]]) as usize; - offset += 2; + let hs_len = reader.read_u16_le()? as usize; - if payload.len() < offset + hs_len { - return Err(Error::MessageTooShort { - expected: offset + hs_len, - got: payload.len(), - }); - } - let handshake_payload = payload[offset..offset + hs_len].to_vec(); + let handshake_payload = reader.read_bytes(hs_len)?.to_vec(); Ok(Self { src_coords, @@ -774,37 +756,18 @@ impl SessionAck { /// Decode from wire format (after 4-byte FSP prefix has been consumed). pub fn decode(payload: &[u8]) -> Result { - if payload.is_empty() { - return Err(Error::MessageTooShort { - expected: 1, - got: 0, - }); - } - let flags = payload[0]; - let mut offset = 1; + let mut reader = Reader::new(payload); + let flags = reader.read_u8()?; - let (src_coords, consumed) = decode_coords(&payload[offset..])?; - offset += consumed; + let (src_coords, consumed) = decode_coords(reader.rest())?; + reader.advance(consumed); - let (dest_coords, consumed) = decode_coords(&payload[offset..])?; - offset += consumed; + let (dest_coords, consumed) = decode_coords(reader.rest())?; + reader.advance(consumed); - if payload.len() < offset + 2 { - return Err(Error::MessageTooShort { - expected: offset + 2, - got: payload.len(), - }); - } - let hs_len = u16::from_le_bytes([payload[offset], payload[offset + 1]]) as usize; - offset += 2; + let hs_len = reader.read_u16_le()? as usize; - if payload.len() < offset + hs_len { - return Err(Error::MessageTooShort { - expected: offset + hs_len, - got: payload.len(), - }); - } - let handshake_payload = payload[offset..offset + hs_len].to_vec(); + let handshake_payload = reader.read_bytes(hs_len)?.to_vec(); Ok(Self { src_coords, @@ -855,49 +818,31 @@ impl SessionMsg3 { /// where ver_phase = 0x03 (version 0, phase MSG3). pub fn encode(&self) -> Vec { // Build body first to compute payload_len - let mut body = Vec::new(); - body.push(self.flags); + let mut body = Writer::new(); + body.write_u8(self.flags); let hs_len = self.handshake_payload.len() as u16; - body.extend_from_slice(&hs_len.to_le_bytes()); - body.extend_from_slice(&self.handshake_payload); + body.write_u16_le(hs_len); + body.write_bytes(&self.handshake_payload); // Prepend 4-byte FSP common prefix let payload_len = body.len() as u16; - let mut buf = Vec::with_capacity(4 + body.len()); - buf.push(0x03); // version 0, phase 0x3 (MSG3) - buf.push(0x00); // flags (must be zero for handshake) - buf.extend_from_slice(&payload_len.to_le_bytes()); - buf.extend_from_slice(&body); - buf + let body = body.into_vec(); + let mut w = Writer::with_capacity(4 + body.len()); + w.write_u8(0x03); // version 0, phase 0x3 (MSG3) + w.write_u8(0x00); // flags (must be zero for handshake) + w.write_u16_le(payload_len); + w.write_bytes(&body); + w.into_vec() } /// Decode from wire format (after 4-byte FSP prefix has been consumed). pub fn decode(payload: &[u8]) -> Result { - if payload.is_empty() { - return Err(Error::MessageTooShort { - expected: 1, - got: 0, - }); - } - let flags = payload[0]; - let mut offset = 1; + let mut reader = Reader::new(payload); + let flags = reader.read_u8()?; - if payload.len() < offset + 2 { - return Err(Error::MessageTooShort { - expected: offset + 2, - got: payload.len(), - }); - } - let hs_len = u16::from_le_bytes([payload[offset], payload[offset + 1]]) as usize; - offset += 2; + let hs_len = reader.read_u16_le()? as usize; - if payload.len() < offset + hs_len { - return Err(Error::MessageTooShort { - expected: offset + hs_len, - got: payload.len(), - }); - } - let handshake_payload = payload[offset..offset + hs_len].to_vec(); + let handshake_payload = reader.read_bytes(hs_len)?.to_vec(); Ok(Self { flags, diff --git a/src/proto/mmp/wire.rs b/src/proto/mmp/wire.rs index 342d407..282f8fa 100644 --- a/src/proto/mmp/wire.rs +++ b/src/proto/mmp/wire.rs @@ -7,6 +7,7 @@ //! between the two layers. Wire format follows the MMP design doc. use crate::proto::Error; +use crate::proto::codec::{Reader, Writer}; // ============================================================================ // SenderReport (msg_type 0x01, 48-byte body including type byte) @@ -82,39 +83,35 @@ pub struct ReceiverReport { impl SenderReport { /// Encode to wire format (48 bytes: msg_type + 3 reserved + 44 payload). pub fn encode(&self) -> Vec { - let mut buf = Vec::with_capacity(48); - buf.push(0x01); // msg_type - buf.extend_from_slice(&[0u8; 3]); // reserved - buf.extend_from_slice(&self.interval_start_counter.to_le_bytes()); - buf.extend_from_slice(&self.interval_end_counter.to_le_bytes()); - buf.extend_from_slice(&self.interval_start_timestamp.to_le_bytes()); - buf.extend_from_slice(&self.interval_end_timestamp.to_le_bytes()); - buf.extend_from_slice(&self.interval_bytes_sent.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_packets_sent.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_bytes_sent.to_le_bytes()); - buf + let mut w = Writer::with_capacity(48); + w.write_u8(0x01); // msg_type + w.write_bytes(&[0u8; 3]); // reserved + w.write_u64_le(self.interval_start_counter); + w.write_u64_le(self.interval_end_counter); + w.write_u32_le(self.interval_start_timestamp); + w.write_u32_le(self.interval_end_timestamp); + w.write_u32_le(self.interval_bytes_sent); + w.write_u64_le(self.cumulative_packets_sent); + w.write_u64_le(self.cumulative_bytes_sent); + w.into_vec() } /// Decode from payload after msg_type byte has been consumed. /// /// `payload` starts at the reserved bytes (offset 1 in the wire format). pub fn decode(payload: &[u8]) -> Result { - if payload.len() < 47 { - return Err(Error::MessageTooShort { - expected: 47, - got: payload.len(), - }); - } + let mut reader = Reader::new(payload); + reader.require(47)?; // Skip 3 reserved bytes - let p = &payload[3..]; + reader.advance(3); Ok(Self { - interval_start_counter: u64::from_le_bytes(p[0..8].try_into().unwrap()), - interval_end_counter: u64::from_le_bytes(p[8..16].try_into().unwrap()), - interval_start_timestamp: u32::from_le_bytes(p[16..20].try_into().unwrap()), - interval_end_timestamp: u32::from_le_bytes(p[20..24].try_into().unwrap()), - interval_bytes_sent: u32::from_le_bytes(p[24..28].try_into().unwrap()), - cumulative_packets_sent: u64::from_le_bytes(p[28..36].try_into().unwrap()), - cumulative_bytes_sent: u64::from_le_bytes(p[36..44].try_into().unwrap()), + interval_start_counter: reader.read_u64_le()?, + interval_end_counter: reader.read_u64_le()?, + interval_start_timestamp: reader.read_u32_le()?, + interval_end_timestamp: reader.read_u32_le()?, + interval_bytes_sent: reader.read_u32_le()?, + cumulative_packets_sent: reader.read_u64_le()?, + cumulative_bytes_sent: reader.read_u64_le()?, }) } } @@ -122,55 +119,54 @@ impl SenderReport { impl ReceiverReport { /// Encode to wire format (68 bytes: msg_type + 3 reserved + 64 payload). pub fn encode(&self) -> Vec { - let mut buf = Vec::with_capacity(68); - buf.push(0x02); // msg_type - buf.extend_from_slice(&[0u8; 3]); // reserved - buf.extend_from_slice(&self.highest_counter.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_packets_recv.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_bytes_recv.to_le_bytes()); - buf.extend_from_slice(&self.timestamp_echo.to_le_bytes()); - buf.extend_from_slice(&self.dwell_time.to_le_bytes()); - buf.extend_from_slice(&self.max_burst_loss.to_le_bytes()); - buf.extend_from_slice(&self.mean_burst_loss.to_le_bytes()); - buf.extend_from_slice(&[0u8; 2]); // reserved - buf.extend_from_slice(&self.jitter.to_le_bytes()); - buf.extend_from_slice(&self.ecn_ce_count.to_le_bytes()); - buf.extend_from_slice(&self.owd_trend.to_le_bytes()); - buf.extend_from_slice(&self.burst_loss_count.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_reorder_count.to_le_bytes()); - buf.extend_from_slice(&self.interval_packets_recv.to_le_bytes()); - buf.extend_from_slice(&self.interval_bytes_recv.to_le_bytes()); - buf + let mut w = Writer::with_capacity(68); + w.write_u8(0x02); // msg_type + w.write_bytes(&[0u8; 3]); // reserved + w.write_u64_le(self.highest_counter); + w.write_u64_le(self.cumulative_packets_recv); + w.write_u64_le(self.cumulative_bytes_recv); + w.write_u32_le(self.timestamp_echo); + w.write_u16_le(self.dwell_time); + w.write_u16_le(self.max_burst_loss); + w.write_u16_le(self.mean_burst_loss); + w.write_bytes(&[0u8; 2]); // reserved + w.write_u32_le(self.jitter); + w.write_u32_le(self.ecn_ce_count); + w.write_bytes(&self.owd_trend.to_le_bytes()); + w.write_u32_le(self.burst_loss_count); + w.write_u32_le(self.cumulative_reorder_count); + w.write_u32_le(self.interval_packets_recv); + w.write_u32_le(self.interval_bytes_recv); + w.into_vec() } /// Decode from payload after msg_type byte has been consumed. /// /// `payload` starts at the reserved bytes (offset 1 in the wire format). pub fn decode(payload: &[u8]) -> Result { - if payload.len() < 67 { - return Err(Error::MessageTooShort { - expected: 67, - got: payload.len(), - }); - } + let mut reader = Reader::new(payload); + reader.require(67)?; // Skip 3 reserved bytes - let p = &payload[3..]; + reader.advance(3); Ok(Self { - highest_counter: u64::from_le_bytes(p[0..8].try_into().unwrap()), - cumulative_packets_recv: u64::from_le_bytes(p[8..16].try_into().unwrap()), - cumulative_bytes_recv: u64::from_le_bytes(p[16..24].try_into().unwrap()), - timestamp_echo: u32::from_le_bytes(p[24..28].try_into().unwrap()), - dwell_time: u16::from_le_bytes(p[28..30].try_into().unwrap()), - max_burst_loss: u16::from_le_bytes(p[30..32].try_into().unwrap()), - mean_burst_loss: u16::from_le_bytes(p[32..34].try_into().unwrap()), + highest_counter: reader.read_u64_le()?, + cumulative_packets_recv: reader.read_u64_le()?, + cumulative_bytes_recv: reader.read_u64_le()?, + timestamp_echo: reader.read_u32_le()?, + dwell_time: reader.read_u16_le()?, + max_burst_loss: reader.read_u16_le()?, + mean_burst_loss: reader.read_u16_le()?, // skip 2 reserved bytes at p[34..36] - jitter: u32::from_le_bytes(p[36..40].try_into().unwrap()), - ecn_ce_count: u32::from_le_bytes(p[40..44].try_into().unwrap()), - owd_trend: i32::from_le_bytes(p[44..48].try_into().unwrap()), - burst_loss_count: u32::from_le_bytes(p[48..52].try_into().unwrap()), - cumulative_reorder_count: u32::from_le_bytes(p[52..56].try_into().unwrap()), - interval_packets_recv: u32::from_le_bytes(p[56..60].try_into().unwrap()), - interval_bytes_recv: u32::from_le_bytes(p[60..64].try_into().unwrap()), + jitter: { + reader.advance(2); + reader.read_u32_le()? + }, + ecn_ce_count: reader.read_u32_le()?, + owd_trend: i32::from_le_bytes(reader.read_array::<4>()?), + burst_loss_count: reader.read_u32_le()?, + cumulative_reorder_count: reader.read_u32_le()?, + interval_packets_recv: reader.read_u32_le()?, + interval_bytes_recv: reader.read_u32_le()?, }) } } @@ -214,36 +210,32 @@ pub const SESSION_SENDER_REPORT_SIZE: usize = 46; impl SessionSenderReport { /// Encode to wire format (46 bytes body). pub fn encode(&self) -> Vec { - let mut buf = Vec::with_capacity(SESSION_SENDER_REPORT_SIZE); - buf.extend_from_slice(&[0u8; 2]); // reserved - buf.extend_from_slice(&self.interval_start_counter.to_le_bytes()); - buf.extend_from_slice(&self.interval_end_counter.to_le_bytes()); - buf.extend_from_slice(&self.interval_start_timestamp.to_le_bytes()); - buf.extend_from_slice(&self.interval_end_timestamp.to_le_bytes()); - buf.extend_from_slice(&self.interval_bytes_sent.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_packets_sent.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_bytes_sent.to_le_bytes()); - buf + let mut w = Writer::with_capacity(SESSION_SENDER_REPORT_SIZE); + w.write_bytes(&[0u8; 2]); // reserved + w.write_u64_le(self.interval_start_counter); + w.write_u64_le(self.interval_end_counter); + w.write_u32_le(self.interval_start_timestamp); + w.write_u32_le(self.interval_end_timestamp); + w.write_u32_le(self.interval_bytes_sent); + w.write_u64_le(self.cumulative_packets_sent); + w.write_u64_le(self.cumulative_bytes_sent); + w.into_vec() } /// Decode from body (after FSP inner header has been stripped). pub fn decode(body: &[u8]) -> Result { - if body.len() < SESSION_SENDER_REPORT_SIZE { - return Err(Error::MessageTooShort { - expected: SESSION_SENDER_REPORT_SIZE, - got: body.len(), - }); - } + let mut reader = Reader::new(body); + reader.require(SESSION_SENDER_REPORT_SIZE)?; // Skip 2 reserved bytes - let p = &body[2..]; + reader.advance(2); Ok(Self { - interval_start_counter: u64::from_le_bytes(p[0..8].try_into().unwrap()), - interval_end_counter: u64::from_le_bytes(p[8..16].try_into().unwrap()), - interval_start_timestamp: u32::from_le_bytes(p[16..20].try_into().unwrap()), - interval_end_timestamp: u32::from_le_bytes(p[20..24].try_into().unwrap()), - interval_bytes_sent: u32::from_le_bytes(p[24..28].try_into().unwrap()), - cumulative_packets_sent: u64::from_le_bytes(p[28..36].try_into().unwrap()), - cumulative_bytes_sent: u64::from_le_bytes(p[36..44].try_into().unwrap()), + interval_start_counter: reader.read_u64_le()?, + interval_end_counter: reader.read_u64_le()?, + interval_start_timestamp: reader.read_u32_le()?, + interval_end_timestamp: reader.read_u32_le()?, + interval_bytes_sent: reader.read_u32_le()?, + cumulative_packets_sent: reader.read_u64_le()?, + cumulative_bytes_sent: reader.read_u64_le()?, }) } } @@ -297,52 +289,51 @@ pub const SESSION_RECEIVER_REPORT_SIZE: usize = 66; impl SessionReceiverReport { /// Encode to wire format (66 bytes body). pub fn encode(&self) -> Vec { - let mut buf = Vec::with_capacity(SESSION_RECEIVER_REPORT_SIZE); - buf.extend_from_slice(&[0u8; 2]); // reserved - buf.extend_from_slice(&self.highest_counter.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_packets_recv.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_bytes_recv.to_le_bytes()); - buf.extend_from_slice(&self.timestamp_echo.to_le_bytes()); - buf.extend_from_slice(&self.dwell_time.to_le_bytes()); - buf.extend_from_slice(&self.max_burst_loss.to_le_bytes()); - buf.extend_from_slice(&self.mean_burst_loss.to_le_bytes()); - buf.extend_from_slice(&[0u8; 2]); // reserved - buf.extend_from_slice(&self.jitter.to_le_bytes()); - buf.extend_from_slice(&self.ecn_ce_count.to_le_bytes()); - buf.extend_from_slice(&self.owd_trend.to_le_bytes()); - buf.extend_from_slice(&self.burst_loss_count.to_le_bytes()); - buf.extend_from_slice(&self.cumulative_reorder_count.to_le_bytes()); - buf.extend_from_slice(&self.interval_packets_recv.to_le_bytes()); - buf.extend_from_slice(&self.interval_bytes_recv.to_le_bytes()); - buf + let mut w = Writer::with_capacity(SESSION_RECEIVER_REPORT_SIZE); + w.write_bytes(&[0u8; 2]); // reserved + w.write_u64_le(self.highest_counter); + w.write_u64_le(self.cumulative_packets_recv); + w.write_u64_le(self.cumulative_bytes_recv); + w.write_u32_le(self.timestamp_echo); + w.write_u16_le(self.dwell_time); + w.write_u16_le(self.max_burst_loss); + w.write_u16_le(self.mean_burst_loss); + w.write_bytes(&[0u8; 2]); // reserved + w.write_u32_le(self.jitter); + w.write_u32_le(self.ecn_ce_count); + w.write_bytes(&self.owd_trend.to_le_bytes()); + w.write_u32_le(self.burst_loss_count); + w.write_u32_le(self.cumulative_reorder_count); + w.write_u32_le(self.interval_packets_recv); + w.write_u32_le(self.interval_bytes_recv); + w.into_vec() } /// Decode from body (after FSP inner header has been stripped). pub fn decode(body: &[u8]) -> Result { - if body.len() < SESSION_RECEIVER_REPORT_SIZE { - return Err(Error::MessageTooShort { - expected: SESSION_RECEIVER_REPORT_SIZE, - got: body.len(), - }); - } + let mut reader = Reader::new(body); + reader.require(SESSION_RECEIVER_REPORT_SIZE)?; // Skip 2 reserved bytes - let p = &body[2..]; + reader.advance(2); Ok(Self { - highest_counter: u64::from_le_bytes(p[0..8].try_into().unwrap()), - cumulative_packets_recv: u64::from_le_bytes(p[8..16].try_into().unwrap()), - cumulative_bytes_recv: u64::from_le_bytes(p[16..24].try_into().unwrap()), - timestamp_echo: u32::from_le_bytes(p[24..28].try_into().unwrap()), - dwell_time: u16::from_le_bytes(p[28..30].try_into().unwrap()), - max_burst_loss: u16::from_le_bytes(p[30..32].try_into().unwrap()), - mean_burst_loss: u16::from_le_bytes(p[32..34].try_into().unwrap()), + highest_counter: reader.read_u64_le()?, + cumulative_packets_recv: reader.read_u64_le()?, + cumulative_bytes_recv: reader.read_u64_le()?, + timestamp_echo: reader.read_u32_le()?, + dwell_time: reader.read_u16_le()?, + max_burst_loss: reader.read_u16_le()?, + mean_burst_loss: reader.read_u16_le()?, // skip 2 reserved bytes at p[34..36] - jitter: u32::from_le_bytes(p[36..40].try_into().unwrap()), - ecn_ce_count: u32::from_le_bytes(p[40..44].try_into().unwrap()), - owd_trend: i32::from_le_bytes(p[44..48].try_into().unwrap()), - burst_loss_count: u32::from_le_bytes(p[48..52].try_into().unwrap()), - cumulative_reorder_count: u32::from_le_bytes(p[52..56].try_into().unwrap()), - interval_packets_recv: u32::from_le_bytes(p[56..60].try_into().unwrap()), - interval_bytes_recv: u32::from_le_bytes(p[60..64].try_into().unwrap()), + jitter: { + reader.advance(2); + reader.read_u32_le()? + }, + ecn_ce_count: reader.read_u32_le()?, + owd_trend: i32::from_le_bytes(reader.read_array::<4>()?), + burst_loss_count: reader.read_u32_le()?, + cumulative_reorder_count: reader.read_u32_le()?, + interval_packets_recv: reader.read_u32_le()?, + interval_bytes_recv: reader.read_u32_le()?, }) } } @@ -380,14 +371,10 @@ impl PathMtuNotification { /// Decode from body (after FSP inner header has been stripped). pub fn decode(body: &[u8]) -> Result { - if body.len() < PATH_MTU_NOTIFICATION_SIZE { - return Err(Error::MessageTooShort { - expected: PATH_MTU_NOTIFICATION_SIZE, - got: body.len(), - }); - } + let mut reader = Reader::new(body); + reader.require(PATH_MTU_NOTIFICATION_SIZE)?; Ok(Self { - path_mtu: u16::from_le_bytes([body[0], body[1]]), + path_mtu: reader.read_u16_le()?, }) } } diff --git a/src/proto/mod.rs b/src/proto/mod.rs index e5fb961..b48e7b7 100644 --- a/src/proto/mod.rs +++ b/src/proto/mod.rs @@ -7,6 +7,7 @@ mod error; pub use error::Error; pub(crate) mod bloom; +pub(crate) mod codec; pub(crate) mod coord; pub(crate) mod discovery; pub(crate) mod fmp; diff --git a/src/proto/routing/wire.rs b/src/proto/routing/wire.rs index 803c90c..c867629 100644 --- a/src/proto/routing/wire.rs +++ b/src/proto/routing/wire.rs @@ -12,6 +12,7 @@ use crate::NodeAddr; use crate::proto::Error; +use crate::proto::codec::{Reader, Writer}; use crate::proto::stp::{ TreeCoordinate, decode_optional_coords, encode_coords, encode_empty_coords, }; @@ -109,34 +110,29 @@ impl CoordsRequired { pub fn encode(&self) -> Vec { // Body: msg_type + flags(reserved) + dest_addr + reporter let body_len = 1 + 1 + 16 + 16; // 34 bytes - let mut buf = Vec::with_capacity(4 + body_len); + let mut w = Writer::with_capacity(4 + body_len); // FSP prefix: version 0, phase 0x0, U flag set - buf.push(0x00); // version 0, phase 0x0 - buf.push(0x04); // U flag + w.write_u8(0x00); // version 0, phase 0x0 + w.write_u8(0x04); // U flag let payload_len = body_len as u16; - buf.extend_from_slice(&payload_len.to_le_bytes()); + w.write_u16_le(payload_len); // msg_type byte (after prefix, before body) - buf.push(RoutingSignalType::CoordsRequired.to_byte()); - buf.push(0x00); // reserved flags - buf.extend_from_slice(self.dest_addr.as_bytes()); - buf.extend_from_slice(self.reporter.as_bytes()); - buf + w.write_u8(RoutingSignalType::CoordsRequired.to_byte()); + w.write_u8(0x00); // reserved flags + w.write_bytes(self.dest_addr.as_bytes()); + w.write_bytes(self.reporter.as_bytes()); + w.into_vec() } /// Decode from wire format (after FSP prefix and msg_type byte consumed). pub fn decode(payload: &[u8]) -> Result { // flags(1) + dest_addr(16) + reporter(16) = 33 - if payload.len() < 33 { - return Err(Error::MessageTooShort { - expected: 33, - got: payload.len(), - }); - } + let mut reader = Reader::new(payload); + reader.require(33)?; // payload[0] is flags (reserved, ignored) - let mut dest_bytes = [0u8; 16]; - dest_bytes.copy_from_slice(&payload[1..17]); - let mut reporter_bytes = [0u8; 16]; - reporter_bytes.copy_from_slice(&payload[17..33]); + reader.advance(1); + let dest_bytes = reader.read_array::<16>()?; + let reporter_bytes = reader.read_array::<16>()?; Ok(Self { dest_addr: NodeAddr::from_bytes(dest_bytes), @@ -217,19 +213,14 @@ impl PathBroken { /// Decode from wire format (after FSP prefix and msg_type byte consumed). pub fn decode(payload: &[u8]) -> Result { // flags(1) + dest_addr(16) + reporter(16) + coords_count(2) = 35 minimum - if payload.len() < 35 { - return Err(Error::MessageTooShort { - expected: 35, - got: payload.len(), - }); - } + let mut reader = Reader::new(payload); + reader.require(35)?; // payload[0] is flags (reserved, ignored) - let mut dest_bytes = [0u8; 16]; - dest_bytes.copy_from_slice(&payload[1..17]); - let mut reporter_bytes = [0u8; 16]; - reporter_bytes.copy_from_slice(&payload[17..33]); + reader.advance(1); + let dest_bytes = reader.read_array::<16>()?; + let reporter_bytes = reader.read_array::<16>()?; - let (last_known_coords, _consumed) = decode_optional_coords(&payload[33..])?; + let (last_known_coords, _consumed) = decode_optional_coords(reader.rest())?; Ok(Self { dest_addr: NodeAddr::from_bytes(dest_bytes), @@ -284,36 +275,31 @@ impl MtuExceeded { /// Error signals use phase=0x0 with U flag set. pub fn encode(&self) -> Vec { let body_len = MTU_EXCEEDED_SIZE; // 36 bytes - let mut buf = Vec::with_capacity(4 + body_len); + let mut w = Writer::with_capacity(4 + body_len); // FSP prefix: version 0, phase 0x0, U flag set - buf.push(0x00); // version 0, phase 0x0 - buf.push(0x04); // U flag + w.write_u8(0x00); // version 0, phase 0x0 + w.write_u8(0x04); // U flag let payload_len = body_len as u16; - buf.extend_from_slice(&payload_len.to_le_bytes()); + w.write_u16_le(payload_len); // msg_type byte - buf.push(RoutingSignalType::MtuExceeded.to_byte()); - buf.push(0x00); // reserved flags - buf.extend_from_slice(self.dest_addr.as_bytes()); - buf.extend_from_slice(self.reporter.as_bytes()); - buf.extend_from_slice(&self.mtu.to_le_bytes()); - buf + w.write_u8(RoutingSignalType::MtuExceeded.to_byte()); + w.write_u8(0x00); // reserved flags + w.write_bytes(self.dest_addr.as_bytes()); + w.write_bytes(self.reporter.as_bytes()); + w.write_u16_le(self.mtu); + w.into_vec() } /// Decode from wire format (after FSP prefix and msg_type byte consumed). pub fn decode(payload: &[u8]) -> Result { // flags(1) + dest_addr(16) + reporter(16) + mtu(2) = 35 - if payload.len() < 35 { - return Err(Error::MessageTooShort { - expected: 35, - got: payload.len(), - }); - } + let mut reader = Reader::new(payload); + reader.require(35)?; // payload[0] is flags (reserved, ignored) - let mut dest_bytes = [0u8; 16]; - dest_bytes.copy_from_slice(&payload[1..17]); - let mut reporter_bytes = [0u8; 16]; - reporter_bytes.copy_from_slice(&payload[17..33]); - let mtu = u16::from_le_bytes([payload[33], payload[34]]); + reader.advance(1); + let dest_bytes = reader.read_array::<16>()?; + let reporter_bytes = reader.read_array::<16>()?; + let mtu = reader.read_u16_le()?; Ok(Self { dest_addr: NodeAddr::from_bytes(dest_bytes), diff --git a/src/proto/stp/wire.rs b/src/proto/stp/wire.rs index 45a9efb..805bc00 100644 --- a/src/proto/stp/wire.rs +++ b/src/proto/stp/wire.rs @@ -3,6 +3,7 @@ use super::{CoordEntry, ParentDeclaration, TreeCoordinate, TreeError}; use crate::NodeAddr; use crate::proto::Error; +use crate::proto::codec::{Reader, Writer}; use crate::proto::link::LinkMessageType; use secp256k1::schnorr::Signature; @@ -103,121 +104,72 @@ impl TreeAnnounce { let entries = self.ancestry.entries(); let ancestry_count = entries.len() as u16; let size = 1 + Self::MIN_PAYLOAD_SIZE + entries.len() * CoordEntry::WIRE_SIZE; - let mut buf = Vec::with_capacity(size); + let mut w = Writer::with_capacity(size); // msg_type - buf.push(LinkMessageType::TreeAnnounce.to_byte()); + w.write_u8(LinkMessageType::TreeAnnounce.to_byte()); // version - buf.push(Self::VERSION_1); + w.write_u8(Self::VERSION_1); // sequence (8 LE) - buf.extend_from_slice(&self.declaration.sequence().to_le_bytes()); + w.write_u64_le(self.declaration.sequence()); // timestamp (8 LE) - buf.extend_from_slice(&self.declaration.timestamp().to_le_bytes()); + w.write_u64_le(self.declaration.timestamp()); // parent (16) - buf.extend_from_slice(self.declaration.parent_id().as_bytes()); + w.write_bytes(self.declaration.parent_id().as_bytes()); // ancestry_count (2 LE) - buf.extend_from_slice(&ancestry_count.to_le_bytes()); + w.write_u16_le(ancestry_count); // ancestry entries (32 bytes each) for entry in entries { - buf.extend_from_slice(entry.node_addr.as_bytes()); // 16 - buf.extend_from_slice(&entry.sequence.to_le_bytes()); // 8 - buf.extend_from_slice(&entry.timestamp.to_le_bytes()); // 8 + w.write_bytes(entry.node_addr.as_bytes()); // 16 + w.write_u64_le(entry.sequence); // 8 + w.write_u64_le(entry.timestamp); // 8 } // outer signature (64) - buf.extend_from_slice(signature.as_ref()); + w.write_bytes(signature.as_ref()); - Ok(buf) + Ok(w.into_vec()) } /// Decode from link-layer payload (after msg_type byte stripped by dispatcher). /// /// The payload starts with the version byte. pub fn decode(payload: &[u8]) -> Result { - if payload.len() < Self::MIN_PAYLOAD_SIZE { - return Err(Error::MessageTooShort { - expected: Self::MIN_PAYLOAD_SIZE, - got: payload.len(), - }); - } - - let mut pos = 0; + let mut reader = Reader::new(payload); + reader.require(Self::MIN_PAYLOAD_SIZE)?; // version - let version = payload[pos]; - pos += 1; + let version = reader.read_u8()?; if version != Self::VERSION_1 { return Err(Error::UnsupportedVersion(version)); } // sequence (8 LE) - let sequence = u64::from_le_bytes( - payload[pos..pos + 8] - .try_into() - .map_err(|_| Error::Malformed("bad sequence"))?, - ); - pos += 8; + let sequence = reader.read_u64_le()?; // timestamp (8 LE) - let timestamp = u64::from_le_bytes( - payload[pos..pos + 8] - .try_into() - .map_err(|_| Error::Malformed("bad timestamp"))?, - ); - pos += 8; + let timestamp = reader.read_u64_le()?; // parent (16) - let parent = NodeAddr::from_bytes( - payload[pos..pos + 16] - .try_into() - .map_err(|_| Error::Malformed("bad parent"))?, - ); - pos += 16; + let parent = NodeAddr::from_bytes(reader.read_array::<16>()?); // ancestry_count (2 LE) - let ancestry_count = u16::from_le_bytes( - payload[pos..pos + 2] - .try_into() - .map_err(|_| Error::Malformed("bad ancestry count"))?, - ) as usize; - pos += 2; + let ancestry_count = reader.read_u16_le()? as usize; // Validate remaining length: entries + signature let expected_remaining = ancestry_count * CoordEntry::WIRE_SIZE + 64; - if payload.len() - pos < expected_remaining { - return Err(Error::MessageTooShort { - expected: pos + expected_remaining, - got: payload.len(), - }); - } + reader.require(expected_remaining)?; // ancestry entries (32 bytes each) let mut entries = Vec::with_capacity(ancestry_count); for _ in 0..ancestry_count { - let node_addr = NodeAddr::from_bytes( - payload[pos..pos + 16] - .try_into() - .map_err(|_| Error::Malformed("bad entry node_addr"))?, - ); - pos += 16; - let entry_seq = u64::from_le_bytes( - payload[pos..pos + 8] - .try_into() - .map_err(|_| Error::Malformed("bad entry sequence"))?, - ); - pos += 8; - let entry_ts = u64::from_le_bytes( - payload[pos..pos + 8] - .try_into() - .map_err(|_| Error::Malformed("bad entry timestamp"))?, - ); - pos += 8; + let node_addr = NodeAddr::from_bytes(reader.read_array::<16>()?); + let entry_seq = reader.read_u64_le()?; + let entry_ts = reader.read_u64_le()?; entries.push(CoordEntry::new(node_addr, entry_seq, entry_ts)); } // signature (64) - let sig_bytes: [u8; 64] = payload[pos..pos + 64] - .try_into() - .map_err(|_| Error::Malformed("bad signature"))?; + let sig_bytes: [u8; 64] = reader.read_array::<64>()?; // Validate the signature parses as a well-formed schnorr signature (the // codec's only crypto touch, §11 w2); store the raw bytes so the in-core // declaration carries no `secp256k1` dependency. Actual verification is a