1use crate::decode::{MINIMUM_WIRE_VERSION, SSP2_MAGIC, WIRE_VERSION};
18use crate::error::{DecodeError, Result};
19use crate::model::frame_type;
20
21#[derive(Debug, Clone, PartialEq, Eq)]
24pub struct ScannedMessage {
25 pub message: Vec<u8>,
27 pub excess: usize,
31}
32
33#[derive(Default)]
40pub struct MessageStreamScanner {
41 buffer: Vec<u8>,
42 offset: Option<usize>,
45 complete: bool,
46}
47
48impl MessageStreamScanner {
49 pub fn new() -> Self {
50 Self::default()
51 }
52
53 pub fn started(&self) -> bool {
55 !self.buffer.is_empty()
56 }
57
58 pub fn push(&mut self, chunk: &[u8]) -> Result<Option<ScannedMessage>> {
64 assert!(
65 !self.complete,
66 "MessageStreamScanner: message already complete"
67 );
68 self.buffer.extend_from_slice(chunk);
69 if self.offset.is_none() {
70 if self.buffer.len() < 8 {
71 return Ok(None);
72 }
73 self.check_header()?;
74 self.offset = Some(8);
75 }
76 let mut offset = self.offset.expect("offset set once header parsed");
77 loop {
78 if self.buffer.len() - offset < 5 {
80 self.offset = Some(offset);
81 return Ok(None);
82 }
83 let frame_type = self.buffer[offset];
84 let frame_length = u32::from_le_bytes([
85 self.buffer[offset + 1],
86 self.buffer[offset + 2],
87 self.buffer[offset + 3],
88 self.buffer[offset + 4],
89 ]) as usize;
90 let frame_end = offset + 5 + frame_length;
91 if self.buffer.len() < frame_end {
92 self.offset = Some(offset);
93 return Ok(None);
94 }
95 offset = frame_end;
96 if frame_type == frame_type::END {
97 self.complete = true;
98 self.offset = Some(offset);
99 return Ok(Some(ScannedMessage {
100 message: self.buffer[..frame_end].to_vec(),
101 excess: self.buffer.len() - frame_end,
102 }));
103 }
104 }
105 }
106
107 fn check_header(&self) -> Result<()> {
108 let b = &self.buffer;
109 if b[0..4] != SSP2_MAGIC[..] {
110 return Err(DecodeError::invalid("bad envelope magic (expected SSP2)"));
111 }
112 let wire_version = u16::from_le_bytes([b[4], b[5]]);
113 if !(MINIMUM_WIRE_VERSION..=WIRE_VERSION).contains(&wire_version) {
114 return Err(DecodeError::invalid(format!(
115 "unsupported wireVersion {wire_version}"
116 )));
117 }
118 if b[6] != 0x01 && b[6] != 0x02 {
119 return Err(DecodeError::invalid(format!(
120 "unknown msgKind byte 0x{:02x}",
121 b[6]
122 )));
123 }
124 if b[7] != 0x00 {
125 return Err(DecodeError::invalid(format!(
126 "envelope flags must be 0x00, got 0x{:02x}",
127 b[7]
128 )));
129 }
130 Ok(())
131 }
132}
133
134#[cfg(test)]
135mod tests {
136 use super::*;
137 use crate::encode::encode_message;
138 use crate::model::{Frame, Message, MsgKind};
139
140 fn request_bytes() -> Vec<u8> {
143 let message = Message {
144 wire_version: 1,
145 msg_kind: MsgKind::Request,
146 frames: vec![
147 Frame::ReqHeader {
148 client_id: "c1".to_owned(),
149 schema_version: 1,
150 log_epoch: None,
151 },
152 Frame::PullHeader {
153 limit_commits: 0,
154 limit_snapshot_rows: 0,
155 max_snapshot_pages: 0,
156 accept: 0b0011,
157 },
158 ],
159 };
160 encode_message(&message)
161 }
162
163 #[test]
164 fn whole_message_in_one_chunk_completes_with_zero_excess() {
165 let request = request_bytes();
166 let mut scanner = MessageStreamScanner::new();
167 let result = scanner.push(&request).unwrap().expect("complete");
168 assert_eq!(result.excess, 0);
169 assert_eq!(result.message, request);
170 }
171
172 #[test]
173 fn every_split_point_reassembles_byte_exactly() {
174 let request = request_bytes();
175 for split in 1..request.len() {
176 let mut scanner = MessageStreamScanner::new();
177 assert!(
178 scanner.push(&request[..split]).unwrap().is_none(),
179 "split {split}: first half must be incomplete"
180 );
181 let result = scanner
182 .push(&request[split..])
183 .unwrap()
184 .unwrap_or_else(|| panic!("split {split}: second half must complete"));
185 assert_eq!(result.excess, 0, "split {split}");
186 assert_eq!(result.message, request, "split {split}");
187 }
188 }
189
190 #[test]
191 fn one_byte_at_a_time_trickle_completes() {
192 let request = request_bytes();
193 let mut scanner = MessageStreamScanner::new();
194 let mut result = None;
195 for byte in &request {
196 result = scanner.push(&[*byte]).unwrap();
197 }
198 let result = result.expect("complete after last byte");
199 assert_eq!(result.excess, 0);
200 assert_eq!(result.message, request);
201 }
202
203 #[test]
204 fn bytes_past_end_are_reported_as_excess() {
205 let request = request_bytes();
206 let mut with_excess = request.clone();
207 with_excess.extend_from_slice(&[0xaa, 0xbb, 0xcc]);
208 let mut scanner = MessageStreamScanner::new();
209 let result = scanner.push(&with_excess).unwrap().expect("complete");
210 assert_eq!(result.excess, 3);
211 assert_eq!(result.message, request);
212 }
213
214 #[test]
215 fn two_message_stream_every_split_reassembles_first_exactly() {
216 let first = request_bytes();
222 let second = request_bytes();
223 let mut stream = first.clone();
224 stream.extend_from_slice(&second);
225 for split in 1..stream.len() {
226 let mut scanner = MessageStreamScanner::new();
227 let (result, buffered_at_completion) = match scanner.push(&stream[..split]).unwrap() {
231 Some(done) => (done, split),
232 None => {
233 let done = scanner
234 .push(&stream[split..])
235 .unwrap()
236 .unwrap_or_else(|| panic!("split {split}: first message must complete"));
237 (done, stream.len())
238 }
239 };
240 assert_eq!(result.message, first, "split {split}: first message bytes");
241 assert_eq!(
242 result.excess,
243 buffered_at_completion - first.len(),
244 "split {split}: excess equals second-message bytes buffered at END"
245 );
246 }
247 }
248
249 #[test]
250 fn bad_magic_is_a_decode_error_once_header_arrived() {
251 let mut request = request_bytes();
252 request[0] = 0x58;
253 let mut scanner = MessageStreamScanner::new();
254 assert!(scanner.push(&request[..4]).unwrap().is_none());
255 assert!(scanner.push(&request[4..]).is_err());
256 }
257
258 #[test]
259 fn non_zero_flags_and_unknown_msg_kind_are_decode_errors() {
260 for (index, value) in [(6usize, 0x03u8), (7usize, 0x01u8)] {
261 let mut request = request_bytes();
262 request[index] = value;
263 let mut scanner = MessageStreamScanner::new();
264 assert!(scanner.push(&request).is_err(), "index {index}");
265 }
266 }
267
268 #[test]
269 #[should_panic(expected = "already complete")]
270 fn push_after_completion_panics() {
271 let request = request_bytes();
272 let mut scanner = MessageStreamScanner::new();
273 scanner.push(&request).unwrap();
274 let _ = scanner.push(&[0]);
275 }
276}