1use 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 #[must_use]
21 pub fn new() -> Self {
22 Self {
23 buffer: Vec::with_capacity(8 * 1024),
24 }
25 }
26
27 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 if self.buffer.len() >= HEADER_SIZE {
71 self.validate_pending_header()?;
72 }
73
74 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 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 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 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}