Skip to main content

dcp/transport/
framing.rs

1//! Message framing for DCP protocol.
2//!
3//! Provides length-prefixed framing for reliable message boundaries over TCP streams.
4
5use bytes::{Buf, BufMut, Bytes, BytesMut};
6use thiserror::Error;
7
8/// Protocol version byte
9pub const PROTOCOL_VERSION: u8 = 1;
10
11/// Maximum message size (16MB)
12pub const MAX_MESSAGE_SIZE: u32 = 16 * 1024 * 1024;
13
14/// Frame header size (1 byte version + 4 bytes length)
15pub const FRAME_HEADER_SIZE: usize = 5;
16
17/// Frame header - 5 bytes
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
19pub struct FrameHeader {
20    /// Protocol version (1 byte)
21    pub version: u8,
22    /// Payload length in bytes (4 bytes, big-endian)
23    pub length: u32,
24}
25
26impl FrameHeader {
27    /// Create a new frame header
28    pub fn new(length: u32) -> Self {
29        Self {
30            version: PROTOCOL_VERSION,
31            length,
32        }
33    }
34
35    /// Encode header to bytes
36    pub fn encode(&self, dst: &mut BytesMut) {
37        dst.put_u8(self.version);
38        dst.put_u32(self.length);
39    }
40
41    /// Decode header from bytes
42    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    /// Total frame size (header + payload)
53    pub fn frame_size(&self) -> usize {
54        FRAME_HEADER_SIZE + self.length as usize
55    }
56}
57
58/// Frame errors
59#[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/// Frame codec for encoding/decoding messages
72#[derive(Debug)]
73pub struct FrameCodec {
74    /// Maximum allowed message size
75    max_size: u32,
76    /// Read buffer for partial frames
77    read_buffer: BytesMut,
78}
79
80impl Default for FrameCodec {
81    fn default() -> Self {
82        Self::new()
83    }
84}
85
86impl FrameCodec {
87    /// Create a new frame codec with default max size
88    pub fn new() -> Self {
89        Self {
90            max_size: MAX_MESSAGE_SIZE,
91            read_buffer: BytesMut::with_capacity(8192),
92        }
93    }
94
95    /// Create a frame codec with custom max size
96    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    /// Get the maximum message size
104    pub fn max_size(&self) -> u32 {
105        self.max_size
106    }
107
108    /// Feed data into the codec's internal buffer
109    pub fn feed(&mut self, data: &[u8]) {
110        self.read_buffer.extend_from_slice(data);
111    }
112
113    /// Check if there's enough data for a complete frame
114    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    /// Get the current buffer size
129    pub fn buffer_size(&self) -> usize {
130        self.read_buffer.len()
131    }
132
133    /// Clear the internal buffer
134    pub fn clear(&mut self) {
135        self.read_buffer.clear();
136    }
137
138    /// Decode a frame from the internal buffer
139    /// Returns None if incomplete, Some(bytes) if complete
140    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    /// Decode a frame from the provided buffer
157    /// Returns None if incomplete, Some(bytes) if complete
158    pub fn decode_from(&self, src: &mut BytesMut) -> Result<Option<Bytes>, FrameError> {
159        // Need at least header
160        if src.len() < FRAME_HEADER_SIZE {
161            return Ok(None);
162        }
163
164        // Parse header
165        let header = FrameHeader::decode(src).ok_or(FrameError::Incomplete)?;
166
167        // Validate version
168        if header.version != PROTOCOL_VERSION {
169            return Err(FrameError::InvalidVersion(header.version));
170        }
171
172        // Validate size before allocating
173        if header.length > self.max_size {
174            return Err(FrameError::MessageTooLarge(header.length, self.max_size));
175        }
176
177        // Check if we have complete frame
178        let frame_size = header.frame_size();
179        if src.len() < frame_size {
180            return Ok(None);
181        }
182
183        // Extract payload
184        src.advance(FRAME_HEADER_SIZE);
185        let payload = src.split_to(header.length as usize).freeze();
186
187        Ok(Some(payload))
188    }
189
190    /// Encode a message into a frame
191    pub fn encode(&self, msg: &[u8], dst: &mut BytesMut) -> Result<(), FrameError> {
192        let len = msg.len() as u32;
193
194        // Validate size
195        if len > self.max_size {
196            return Err(FrameError::MessageTooLarge(len, self.max_size));
197        }
198
199        // Reserve space
200        dst.reserve(FRAME_HEADER_SIZE + msg.len());
201
202        // Write header
203        let header = FrameHeader::new(len);
204        header.encode(dst);
205
206        // Write payload
207        dst.extend_from_slice(msg);
208
209        Ok(())
210    }
211
212    /// Encode a message and return the framed bytes
213    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
220/// Convenience function to frame a message
221pub fn frame_message(msg: &[u8]) -> Result<Bytes, FrameError> {
222    FrameCodec::new().encode_to_bytes(msg)
223}
224
225/// Convenience function to unframe a message
226pub 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        // Encode
260        let mut encoded = BytesMut::new();
261        codec.encode(message, &mut encoded).unwrap();
262
263        // Decode
264        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],   // Binary data
276            vec![0xFF; 10000], // Larger binary
277        ];
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        // Encode full message
293        let framed = codec.encode_to_bytes(message).unwrap();
294
295        // Try to decode with only header
296        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        // Try with partial payload
301        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        // Full message should decode
306        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); // Invalid version
325        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        // Encode
339        let framed = codec.encode_to_bytes(message).unwrap();
340
341        // Feed in chunks
342        codec.feed(&framed[..3]); // Partial header
343        assert!(!codec.has_complete_frame());
344
345        codec.feed(&framed[3..FRAME_HEADER_SIZE]); // Rest of header
346        assert!(!codec.has_complete_frame());
347
348        codec.feed(&framed[FRAME_HEADER_SIZE..]); // Payload
349        assert!(codec.has_complete_frame());
350
351        // Decode
352        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        // Encode both
364        let framed1 = codec.encode_to_bytes(msg1).unwrap();
365        let framed2 = codec.encode_to_bytes(msg2).unwrap();
366
367        // Feed both at once
368        codec.feed(&framed1);
369        codec.feed(&framed2);
370
371        // Decode first
372        let decoded1 = codec.decode().unwrap().unwrap();
373        assert_eq!(&decoded1[..], msg1);
374
375        // Decode second
376        let decoded2 = codec.decode().unwrap().unwrap();
377        assert_eq!(&decoded2[..], msg2);
378
379        // No more
380        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]; // Less than header size
396        assert!(FrameHeader::decode(&data).is_none());
397    }
398
399    #[test]
400    fn test_max_message_size_validation() {
401        // Create header claiming huge size
402        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}