Skip to main content

nerve_ipc/
io.rs

1//! Buffered I/O helpers for frame decoding.
2//!
3//! `FrameReader` accumulates bytes from a `Read` and yields
4//! parsed `OwnedFrame`s incrementally, draining consumed bytes.
5
6use crate::codec::decode;
7use crate::constants::{HEADER_SIZE, MAGIC, MAX_PAYLOAD_SIZE, VERSION};
8use crate::error::ProtocolError;
9use crate::frame::OwnedFrame;
10use crate::types::{MessageType, ProtocolErrorKind};
11
12use std::io::Read;
13
14pub struct FrameReader {
15    buffer: Vec<u8>,
16}
17
18impl FrameReader {
19    /// Creates a new `FrameReader` with an 8 KiB internal buffer.
20    #[must_use]
21    pub fn new() -> Self {
22        Self {
23            buffer: Vec::with_capacity(8 * 1024),
24        }
25    }
26
27    /// Read bytes from `reader` and return any newly complete frames.
28    ///
29    /// If the internal buffer already holds partial data from a previous call,
30    /// the reader is drained until the pending frame is complete or EOF.
31    /// If the buffer is empty, exactly one `read` is performed and whatever
32    /// complete frames are present are returned.
33    ///
34    /// Header fields (magic, version, message type, payload length) are
35    /// validated as soon as a full header is in the buffer — before any
36    /// payload bytes are accumulated — so a malicious oversized
37    /// `payload_length` cannot cause a large allocation.
38    ///
39    /// # Errors
40    ///
41    /// | Condition | Error kind |
42    /// |-----------|-----------|
43    /// | Underlying `read()` returns an I/O error | [`ProtocolErrorKind::InternalError`] |
44    /// | Magic bytes mismatch | [`ProtocolErrorKind::InvalidMagic`] |
45    /// | Version mismatch | [`ProtocolErrorKind::UnsupportedVersion`] |
46    /// | Unknown `msg_type` byte | [`ProtocolErrorKind::UnknownMessageType`] |
47    /// | `payload_length > MAX_PAYLOAD_SIZE` | [`ProtocolErrorKind::PayloadTooLarge`] |
48    ///
49    /// # Panics
50    ///
51    /// Never panics.  All slice indexing in the extraction loop is guarded by
52    /// the `buffer.len() - offset < HEADER_SIZE` check, and the fixed-width
53    /// `try_into()` conversions are infallible for the exact slice lengths used.
54    pub fn read_from<R: Read>(&mut self, reader: &mut R) -> Result<Vec<OwnedFrame>, ProtocolError> {
55        let started_with_partial = !self.buffer.is_empty();
56        let mut temp = [0u8; 4096];
57
58        loop {
59            let n = reader
60                .read(&mut temp)
61                .map_err(|_| ProtocolError::new(ProtocolErrorKind::InternalError))?;
62
63            if n > 0 {
64                self.buffer.extend_from_slice(&temp[..n]);
65            }
66
67            // Validate the first pending header the moment we have enough bytes.
68            // This happens BEFORE we know the full payload, preventing large
69            // buffer growth for frames with an invalid or oversized payload_length.
70            if self.buffer.len() >= HEADER_SIZE {
71                self.validate_pending_header()?;
72            }
73
74            // Stop looping when:
75            // - reader returned 0 bytes (EOF / nothing available right now), OR
76            // - we didn't start with partial data (single read per fresh start), OR
77            // - we already have enough bytes to complete the first pending frame.
78            if n == 0 || !started_with_partial || self.has_complete_frame() {
79                break;
80            }
81        }
82
83        let mut frames = Vec::new();
84        let mut offset = 0;
85
86        loop {
87            if self.buffer.len() - offset < HEADER_SIZE {
88                break;
89            }
90
91            // Validate header fields before using payload_length to index into the buffer.
92            let magic = u32::from_le_bytes(self.buffer[offset..offset + 4].try_into().unwrap());
93            if magic != MAGIC {
94                return Err(ProtocolError::new(ProtocolErrorKind::InvalidMagic));
95            }
96            let version =
97                u16::from_le_bytes(self.buffer[offset + 4..offset + 6].try_into().unwrap());
98            if version != VERSION {
99                return Err(ProtocolError::new(ProtocolErrorKind::UnsupportedVersion));
100            }
101            MessageType::try_from(self.buffer[offset + 6])
102                .map_err(|()| ProtocolError::new(ProtocolErrorKind::UnknownMessageType))?;
103
104            let payload_len =
105                u32::from_le_bytes(self.buffer[offset + 16..offset + 20].try_into().unwrap())
106                    as usize;
107            if payload_len > MAX_PAYLOAD_SIZE {
108                return Err(ProtocolError::new(ProtocolErrorKind::PayloadTooLarge));
109            }
110
111            let frame_len = HEADER_SIZE + payload_len;
112            if self.buffer.len() - offset < frame_len {
113                break;
114            }
115
116            let decoded = decode(&self.buffer[offset..offset + frame_len])?;
117            frames.push(OwnedFrame {
118                header: decoded.header,
119                payload: decoded.payload.to_vec(),
120            });
121            offset += frame_len;
122        }
123
124        self.buffer.drain(0..offset);
125        Ok(frames)
126    }
127
128    /// Validate the header at the front of the buffer.
129    ///
130    /// Called only when `buffer.len() >= HEADER_SIZE`.
131    fn validate_pending_header(&self) -> Result<(), ProtocolError> {
132        debug_assert!(self.buffer.len() >= HEADER_SIZE);
133
134        let magic = u32::from_le_bytes(self.buffer[0..4].try_into().unwrap());
135        if magic != MAGIC {
136            return Err(ProtocolError::new(ProtocolErrorKind::InvalidMagic));
137        }
138
139        let version = u16::from_le_bytes(self.buffer[4..6].try_into().unwrap());
140        if version != VERSION {
141            return Err(ProtocolError::new(ProtocolErrorKind::UnsupportedVersion));
142        }
143
144        MessageType::try_from(self.buffer[6])
145            .map_err(|()| ProtocolError::new(ProtocolErrorKind::UnknownMessageType))?;
146
147        let payload_len = u32::from_le_bytes(self.buffer[16..20].try_into().unwrap()) as usize;
148        if payload_len > MAX_PAYLOAD_SIZE {
149            return Err(ProtocolError::new(ProtocolErrorKind::PayloadTooLarge));
150        }
151
152        Ok(())
153    }
154
155    /// Return `true` when the buffer holds at least one complete frame.
156    fn has_complete_frame(&self) -> bool {
157        if self.buffer.len() < HEADER_SIZE {
158            return false;
159        }
160        let payload_len = u32::from_le_bytes(self.buffer[16..20].try_into().unwrap()) as usize;
161        self.buffer.len() >= HEADER_SIZE + payload_len
162    }
163}
164
165impl Default for FrameReader {
166    fn default() -> Self {
167        Self::new()
168    }
169}