Skip to main content

nerve_ipc/
codec.rs

1//! Codec for NERVE frames.
2//!
3//! Provides `encode` and `decode` for the wire format.
4//! Ensures single allocation on encode and enforces limits
5//! like `MAX_PAYLOAD_SIZE`, validating header invariants.
6
7use crate::constants::{MAGIC, MAX_PAYLOAD_SIZE, VERSION};
8use crate::error::ProtocolError;
9use crate::frame::{Frame, FrameHeader};
10use crate::types::{FrameFlags, MessageType, ProtocolErrorKind, RequestId};
11
12/// Encode a frame into a contiguous buffer.
13///
14/// Allocation: exactly one `Vec`.
15///
16/// # Errors
17///
18/// Returns [`ProtocolErrorKind::PayloadTooLarge`] if `payload.len()` exceeds
19/// `MAX_PAYLOAD_SIZE` (1 MiB).
20pub fn encode(
21    msg_type: MessageType,
22    flags: FrameFlags,
23    request_id: RequestId,
24    payload: &[u8],
25) -> Result<Vec<u8>, ProtocolError> {
26    if payload.len() > MAX_PAYLOAD_SIZE {
27        return Err(ProtocolError::new(ProtocolErrorKind::PayloadTooLarge));
28    }
29
30    let mut buf = Vec::with_capacity(FrameHeader::SIZE + payload.len());
31
32    // header (little-endian)
33    buf.extend_from_slice(&MAGIC.to_le_bytes());
34    buf.extend_from_slice(&VERSION.to_le_bytes());
35    buf.push(msg_type as u8);
36    buf.push(flags.bits());
37    buf.extend_from_slice(&request_id.0.to_le_bytes());
38    // payload.len() ≤ MAX_PAYLOAD_SIZE (1 MiB), which fits in u32.
39    buf.extend_from_slice(&u32::try_from(payload.len()).unwrap().to_le_bytes());
40
41    // payload
42    buf.extend_from_slice(payload);
43
44    debug_assert_eq!(buf.len(), FrameHeader::SIZE + payload.len());
45
46    Ok(buf)
47}
48
49/// Decode a frame from a byte slice.
50///
51/// # Errors
52///
53/// | Condition | Error kind |
54/// |-----------|-----------|
55/// | `buffer.len() < 20` | [`ProtocolErrorKind::MalformedFrame`] |
56/// | Magic bytes mismatch | [`ProtocolErrorKind::InvalidMagic`] |
57/// | Version mismatch | [`ProtocolErrorKind::UnsupportedVersion`] |
58/// | Unknown `msg_type` byte | [`ProtocolErrorKind::UnknownMessageType`] |
59/// | `payload_length > MAX_PAYLOAD_SIZE` | [`ProtocolErrorKind::PayloadTooLarge`] |
60/// | Buffer shorter than `20 + payload_length` | [`ProtocolErrorKind::MalformedFrame`] |
61///
62/// # Panics
63///
64/// Never panics.  All slice indexing follows the `buffer.len() >= FrameHeader::SIZE`
65/// guard at the top of the function, and the fixed-width `try_into()` conversions
66/// (e.g. `buffer[0..4].try_into()`) are infallible for the exact slice lengths used.
67pub fn decode(buffer: &[u8]) -> Result<Frame<'_>, ProtocolError> {
68    if buffer.len() < FrameHeader::SIZE {
69        return Err(ProtocolError::new(ProtocolErrorKind::MalformedFrame));
70    }
71
72    // header fields
73    let magic = u32::from_le_bytes(buffer[0..4].try_into().unwrap());
74    if magic != MAGIC {
75        return Err(ProtocolError::new(ProtocolErrorKind::InvalidMagic));
76    }
77
78    let version = u16::from_le_bytes(buffer[4..6].try_into().unwrap());
79    if version != VERSION {
80        return Err(ProtocolError::new(ProtocolErrorKind::UnsupportedVersion));
81    }
82
83    let msg_type = MessageType::try_from(buffer[6])
84        .map_err(|()| ProtocolError::new(ProtocolErrorKind::UnknownMessageType))?;
85
86    let flags = FrameFlags::from_bits_truncate(buffer[7]);
87
88    let request_id = RequestId(u64::from_le_bytes(buffer[8..16].try_into().unwrap()));
89
90    let payload_length_u32 = u32::from_le_bytes(buffer[16..20].try_into().unwrap());
91    let payload_length = payload_length_u32 as usize;
92    if payload_length > MAX_PAYLOAD_SIZE {
93        return Err(ProtocolError::new(ProtocolErrorKind::PayloadTooLarge));
94    }
95
96    let total_len = FrameHeader::SIZE + payload_length;
97    if buffer.len() < total_len {
98        return Err(ProtocolError::new(ProtocolErrorKind::MalformedFrame));
99    }
100
101    let payload = &buffer[FrameHeader::SIZE..total_len];
102
103    Ok(Frame {
104        header: FrameHeader {
105            magic,
106            version,
107            msg_type: msg_type as u8,
108            flags: flags.bits(),
109            request_id: request_id.0,
110            payload_length: payload_length_u32,
111        },
112        payload,
113    })
114}