1use bytes::{Buf, BufMut, Bytes, BytesMut};
6use thiserror::Error;
7
8pub const PROTOCOL_VERSION: u8 = 1;
10
11pub const MAX_MESSAGE_SIZE: u32 = 16 * 1024 * 1024;
13
14pub const FRAME_HEADER_SIZE: usize = 5;
16
17#[derive(Debug, Clone, Copy, PartialEq, Eq)]
19pub struct FrameHeader {
20 pub version: u8,
22 pub length: u32,
24}
25
26impl FrameHeader {
27 pub fn new(length: u32) -> Self {
29 Self {
30 version: PROTOCOL_VERSION,
31 length,
32 }
33 }
34
35 pub fn encode(&self, dst: &mut BytesMut) {
37 dst.put_u8(self.version);
38 dst.put_u32(self.length);
39 }
40
41 pub fn decode(src: &[u8]) -> Option<Self> {
43 if src.len() < FRAME_HEADER_SIZE {
44 return None;
45 }
46 Some(Self {
47 version: src[0],
48 length: u32::from_be_bytes([src[1], src[2], src[3], src[4]]),
49 })
50 }
51
52 pub fn frame_size(&self) -> usize {
54 FRAME_HEADER_SIZE + self.length as usize
55 }
56}
57
58#[derive(Debug, Clone, Error, PartialEq, Eq)]
60pub enum FrameError {
61 #[error("message too large: {0} bytes (max: {1})")]
62 MessageTooLarge(u32, u32),
63 #[error("invalid protocol version: {0}")]
64 InvalidVersion(u8),
65 #[error("incomplete frame")]
66 Incomplete,
67 #[error("empty message")]
68 EmptyMessage,
69}
70
71#[derive(Debug)]
73pub struct FrameCodec {
74 max_size: u32,
76 read_buffer: BytesMut,
78}
79
80impl Default for FrameCodec {
81 fn default() -> Self {
82 Self::new()
83 }
84}
85
86impl FrameCodec {
87 pub fn new() -> Self {
89 Self {
90 max_size: MAX_MESSAGE_SIZE,
91 read_buffer: BytesMut::with_capacity(8192),
92 }
93 }
94
95 pub fn with_max_size(max_size: u32) -> Self {
97 Self {
98 max_size,
99 read_buffer: BytesMut::with_capacity(8192),
100 }
101 }
102
103 pub fn max_size(&self) -> u32 {
105 self.max_size
106 }
107
108 pub fn feed(&mut self, data: &[u8]) {
110 self.read_buffer.extend_from_slice(data);
111 }
112
113 pub fn has_complete_frame(&self) -> bool {
115 if self.read_buffer.len() < FRAME_HEADER_SIZE {
116 return false;
117 }
118 if let Some(header) = FrameHeader::decode(&self.read_buffer) {
119 if header.version != PROTOCOL_VERSION || header.length > self.max_size {
120 return true;
121 }
122 self.read_buffer.len() >= header.frame_size()
123 } else {
124 false
125 }
126 }
127
128 pub fn buffer_size(&self) -> usize {
130 self.read_buffer.len()
131 }
132
133 pub fn clear(&mut self) {
135 self.read_buffer.clear();
136 }
137
138 pub fn decode(&mut self) -> Result<Option<Bytes>, FrameError> {
141 let mut probe = self.read_buffer.clone();
142 match self.decode_from(&mut probe) {
143 Ok(Some(bytes)) => {
144 let consumed = FRAME_HEADER_SIZE + bytes.len();
145 self.read_buffer.advance(consumed);
146 Ok(Some(bytes))
147 }
148 Ok(None) => Ok(None),
149 Err(error) => {
150 self.read_buffer.clear();
151 Err(error)
152 }
153 }
154 }
155
156 pub fn decode_from(&self, src: &mut BytesMut) -> Result<Option<Bytes>, FrameError> {
159 if src.len() < FRAME_HEADER_SIZE {
161 return Ok(None);
162 }
163
164 let header = FrameHeader::decode(src).ok_or(FrameError::Incomplete)?;
166
167 if header.version != PROTOCOL_VERSION {
169 return Err(FrameError::InvalidVersion(header.version));
170 }
171
172 if header.length > self.max_size {
174 return Err(FrameError::MessageTooLarge(header.length, self.max_size));
175 }
176
177 let frame_size = header.frame_size();
179 if src.len() < frame_size {
180 return Ok(None);
181 }
182
183 src.advance(FRAME_HEADER_SIZE);
185 let payload = src.split_to(header.length as usize).freeze();
186
187 Ok(Some(payload))
188 }
189
190 pub fn encode(&self, msg: &[u8], dst: &mut BytesMut) -> Result<(), FrameError> {
192 let len = msg.len() as u32;
193
194 if len > self.max_size {
196 return Err(FrameError::MessageTooLarge(len, self.max_size));
197 }
198
199 dst.reserve(FRAME_HEADER_SIZE + msg.len());
201
202 let header = FrameHeader::new(len);
204 header.encode(dst);
205
206 dst.extend_from_slice(msg);
208
209 Ok(())
210 }
211
212 pub fn encode_to_bytes(&self, msg: &[u8]) -> Result<Bytes, FrameError> {
214 let mut dst = BytesMut::with_capacity(FRAME_HEADER_SIZE + msg.len());
215 self.encode(msg, &mut dst)?;
216 Ok(dst.freeze())
217 }
218}
219
220pub fn frame_message(msg: &[u8]) -> Result<Bytes, FrameError> {
222 FrameCodec::new().encode_to_bytes(msg)
223}
224
225pub fn unframe_message(data: &[u8]) -> Result<Option<Bytes>, FrameError> {
227 let mut src = BytesMut::from(data);
228 FrameCodec::new().decode_from(&mut src)
229}
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234
235 #[test]
236 fn test_frame_header_encode_decode() {
237 let header = FrameHeader::new(1234);
238 let mut buf = BytesMut::new();
239 header.encode(&mut buf);
240
241 assert_eq!(buf.len(), FRAME_HEADER_SIZE);
242
243 let decoded = FrameHeader::decode(&buf).unwrap();
244 assert_eq!(decoded.version, PROTOCOL_VERSION);
245 assert_eq!(decoded.length, 1234);
246 }
247
248 #[test]
249 fn test_frame_header_frame_size() {
250 let header = FrameHeader::new(100);
251 assert_eq!(header.frame_size(), FRAME_HEADER_SIZE + 100);
252 }
253
254 #[test]
255 fn test_frame_codec_encode_decode() {
256 let codec = FrameCodec::new();
257 let message = b"Hello, DCP!";
258
259 let mut encoded = BytesMut::new();
261 codec.encode(message, &mut encoded).unwrap();
262
263 let decoded = codec.decode_from(&mut encoded).unwrap().unwrap();
265 assert_eq!(&decoded[..], message);
266 }
267
268 #[test]
269 fn test_frame_codec_round_trip() {
270 let codec = FrameCodec::new();
271 let messages = vec![
272 b"".to_vec(),
273 b"short".to_vec(),
274 b"a longer message with more content".to_vec(),
275 vec![0u8; 1000], vec![0xFF; 10000], ];
278
279 for msg in messages {
280 let framed = codec.encode_to_bytes(&msg).unwrap();
281 let mut src = BytesMut::from(&framed[..]);
282 let decoded = codec.decode_from(&mut src).unwrap().unwrap();
283 assert_eq!(&decoded[..], &msg[..]);
284 }
285 }
286
287 #[test]
288 fn test_frame_codec_partial_message() {
289 let codec = FrameCodec::new();
290 let message = b"Complete message";
291
292 let framed = codec.encode_to_bytes(message).unwrap();
294
295 let mut partial = BytesMut::from(&framed[..FRAME_HEADER_SIZE]);
297 let result = codec.decode_from(&mut partial).unwrap();
298 assert!(result.is_none());
299
300 let mut partial = BytesMut::from(&framed[..FRAME_HEADER_SIZE + 5]);
302 let result = codec.decode_from(&mut partial).unwrap();
303 assert!(result.is_none());
304
305 let mut full = BytesMut::from(&framed[..]);
307 let result = codec.decode_from(&mut full).unwrap();
308 assert!(result.is_some());
309 assert_eq!(&result.unwrap()[..], message);
310 }
311
312 #[test]
313 fn test_frame_codec_message_too_large() {
314 let codec = FrameCodec::with_max_size(100);
315 let message = vec![0u8; 200];
316
317 let result = codec.encode_to_bytes(&message);
318 assert!(matches!(result, Err(FrameError::MessageTooLarge(200, 100))));
319 }
320
321 #[test]
322 fn test_frame_codec_invalid_version() {
323 let mut data = BytesMut::new();
324 data.put_u8(99); data.put_u32(5);
326 data.extend_from_slice(b"hello");
327
328 let codec = FrameCodec::new();
329 let result = codec.decode_from(&mut data);
330 assert!(matches!(result, Err(FrameError::InvalidVersion(99))));
331 }
332
333 #[test]
334 fn test_frame_codec_feed_and_decode() {
335 let mut codec = FrameCodec::new();
336 let message = b"Test message";
337
338 let framed = codec.encode_to_bytes(message).unwrap();
340
341 codec.feed(&framed[..3]); assert!(!codec.has_complete_frame());
344
345 codec.feed(&framed[3..FRAME_HEADER_SIZE]); assert!(!codec.has_complete_frame());
347
348 codec.feed(&framed[FRAME_HEADER_SIZE..]); assert!(codec.has_complete_frame());
350
351 let decoded = codec.decode().unwrap().unwrap();
353 assert_eq!(&decoded[..], message);
354 assert_eq!(codec.buffer_size(), 0);
355 }
356
357 #[test]
358 fn test_frame_codec_multiple_messages() {
359 let mut codec = FrameCodec::new();
360 let msg1 = b"First";
361 let msg2 = b"Second";
362
363 let framed1 = codec.encode_to_bytes(msg1).unwrap();
365 let framed2 = codec.encode_to_bytes(msg2).unwrap();
366
367 codec.feed(&framed1);
369 codec.feed(&framed2);
370
371 let decoded1 = codec.decode().unwrap().unwrap();
373 assert_eq!(&decoded1[..], msg1);
374
375 let decoded2 = codec.decode().unwrap().unwrap();
377 assert_eq!(&decoded2[..], msg2);
378
379 assert!(!codec.has_complete_frame());
381 }
382
383 #[test]
384 fn test_convenience_functions() {
385 let message = b"Quick test";
386
387 let framed = frame_message(message).unwrap();
388 let unframed = unframe_message(&framed).unwrap().unwrap();
389
390 assert_eq!(&unframed[..], message);
391 }
392
393 #[test]
394 fn test_frame_header_decode_insufficient_data() {
395 let data = [0u8; 3]; assert!(FrameHeader::decode(&data).is_none());
397 }
398
399 #[test]
400 fn test_max_message_size_validation() {
401 let mut data = BytesMut::new();
403 data.put_u8(PROTOCOL_VERSION);
404 data.put_u32(MAX_MESSAGE_SIZE + 1);
405
406 let codec = FrameCodec::new();
407 let result = codec.decode_from(&mut data);
408 assert!(matches!(result, Err(FrameError::MessageTooLarge(_, _))));
409 }
410}