nerve-ipc 0.2.0

Binary framing protocol for local IPC over Unix Domain Sockets
Documentation
//! Buffered I/O helpers for frame decoding.
//!
//! `FrameReader` accumulates bytes from a `Read` and yields
//! parsed `OwnedFrame`s incrementally, draining consumed bytes.

use crate::codec::decode;
use crate::constants::{HEADER_SIZE, MAGIC, MAX_PAYLOAD_SIZE, VERSION};
use crate::error::ProtocolError;
use crate::frame::OwnedFrame;
use crate::types::{MessageType, ProtocolErrorKind};

use std::io::Read;

pub struct FrameReader {
    buffer: Vec<u8>,
}

impl FrameReader {
    /// Creates a new `FrameReader` with an 8 KiB internal buffer.
    #[must_use]
    pub fn new() -> Self {
        Self {
            buffer: Vec::with_capacity(8 * 1024),
        }
    }

    /// Read bytes from `reader` and return any newly complete frames.
    ///
    /// If the internal buffer already holds partial data from a previous call,
    /// the reader is drained until the pending frame is complete or EOF.
    /// If the buffer is empty, exactly one `read` is performed and whatever
    /// complete frames are present are returned.
    ///
    /// Header fields (magic, version, message type, payload length) are
    /// validated as soon as a full header is in the buffer — before any
    /// payload bytes are accumulated — so a malicious oversized
    /// `payload_length` cannot cause a large allocation.
    ///
    /// # Errors
    ///
    /// | Condition | Error kind |
    /// |-----------|-----------|
    /// | Underlying `read()` returns an I/O error | [`ProtocolErrorKind::InternalError`] |
    /// | Magic bytes mismatch | [`ProtocolErrorKind::InvalidMagic`] |
    /// | Version mismatch | [`ProtocolErrorKind::UnsupportedVersion`] |
    /// | Unknown `msg_type` byte | [`ProtocolErrorKind::UnknownMessageType`] |
    /// | `payload_length > MAX_PAYLOAD_SIZE` | [`ProtocolErrorKind::PayloadTooLarge`] |
    ///
    /// # Panics
    ///
    /// Never panics.  All slice indexing in the extraction loop is guarded by
    /// the `buffer.len() - offset < HEADER_SIZE` check, and the fixed-width
    /// `try_into()` conversions are infallible for the exact slice lengths used.
    pub fn read_from<R: Read>(&mut self, reader: &mut R) -> Result<Vec<OwnedFrame>, ProtocolError> {
        let started_with_partial = !self.buffer.is_empty();
        let mut temp = [0u8; 4096];

        loop {
            let n = reader
                .read(&mut temp)
                .map_err(|_| ProtocolError::new(ProtocolErrorKind::InternalError))?;

            if n > 0 {
                self.buffer.extend_from_slice(&temp[..n]);
            }

            // Validate the first pending header the moment we have enough bytes.
            // This happens BEFORE we know the full payload, preventing large
            // buffer growth for frames with an invalid or oversized payload_length.
            if self.buffer.len() >= HEADER_SIZE {
                self.validate_pending_header()?;
            }

            // Stop looping when:
            // - reader returned 0 bytes (EOF / nothing available right now), OR
            // - we didn't start with partial data (single read per fresh start), OR
            // - we already have enough bytes to complete the first pending frame.
            if n == 0 || !started_with_partial || self.has_complete_frame() {
                break;
            }
        }

        let mut frames = Vec::new();
        let mut offset = 0;

        loop {
            if self.buffer.len() - offset < HEADER_SIZE {
                break;
            }

            // Validate header fields before using payload_length to index into the buffer.
            let magic = u32::from_le_bytes(self.buffer[offset..offset + 4].try_into().unwrap());
            if magic != MAGIC {
                return Err(ProtocolError::new(ProtocolErrorKind::InvalidMagic));
            }
            let version =
                u16::from_le_bytes(self.buffer[offset + 4..offset + 6].try_into().unwrap());
            if version != VERSION {
                return Err(ProtocolError::new(ProtocolErrorKind::UnsupportedVersion));
            }
            MessageType::try_from(self.buffer[offset + 6])
                .map_err(|()| ProtocolError::new(ProtocolErrorKind::UnknownMessageType))?;

            let payload_len =
                u32::from_le_bytes(self.buffer[offset + 16..offset + 20].try_into().unwrap())
                    as usize;
            if payload_len > MAX_PAYLOAD_SIZE {
                return Err(ProtocolError::new(ProtocolErrorKind::PayloadTooLarge));
            }

            let frame_len = HEADER_SIZE + payload_len;
            if self.buffer.len() - offset < frame_len {
                break;
            }

            let decoded = decode(&self.buffer[offset..offset + frame_len])?;
            frames.push(OwnedFrame {
                header: decoded.header,
                payload: decoded.payload.to_vec(),
            });
            offset += frame_len;
        }

        self.buffer.drain(0..offset);
        Ok(frames)
    }

    /// Validate the header at the front of the buffer.
    ///
    /// Called only when `buffer.len() >= HEADER_SIZE`.
    fn validate_pending_header(&self) -> Result<(), ProtocolError> {
        debug_assert!(self.buffer.len() >= HEADER_SIZE);

        let magic = u32::from_le_bytes(self.buffer[0..4].try_into().unwrap());
        if magic != MAGIC {
            return Err(ProtocolError::new(ProtocolErrorKind::InvalidMagic));
        }

        let version = u16::from_le_bytes(self.buffer[4..6].try_into().unwrap());
        if version != VERSION {
            return Err(ProtocolError::new(ProtocolErrorKind::UnsupportedVersion));
        }

        MessageType::try_from(self.buffer[6])
            .map_err(|()| ProtocolError::new(ProtocolErrorKind::UnknownMessageType))?;

        let payload_len = u32::from_le_bytes(self.buffer[16..20].try_into().unwrap()) as usize;
        if payload_len > MAX_PAYLOAD_SIZE {
            return Err(ProtocolError::new(ProtocolErrorKind::PayloadTooLarge));
        }

        Ok(())
    }

    /// Return `true` when the buffer holds at least one complete frame.
    fn has_complete_frame(&self) -> bool {
        if self.buffer.len() < HEADER_SIZE {
            return false;
        }
        let payload_len = u32::from_le_bytes(self.buffer[16..20].try_into().unwrap()) as usize;
        self.buffer.len() >= HEADER_SIZE + payload_len
    }
}

impl Default for FrameReader {
    fn default() -> Self {
        Self::new()
    }
}