1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
//! 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()
}
}