diff --git a/src/framing.rs b/src/framing.rs index f1b6fca..be52ff4 100644 --- a/src/framing.rs +++ b/src/framing.rs @@ -181,48 +181,42 @@ impl Frame { } /// Deserialize a frame from a buffer. - /// Returns (frame, remaining_bytes) if successful. + /// Returns (frame, payload_len) if successful. + /// + /// CRITICAL: If the frame is incomplete (partial header or partial payload), + /// NO bytes are consumed from the buffer. The caller retains all buffered + /// bytes and can retry after more data arrives. pub fn deserialize(buf: &mut BytesMut) -> Result<(Self, usize), FrameError> { - // Need at least the header + // Need at least the header — do NOT consume yet if buf.len() < FRAME_HEADER_LEN { return Err(FrameError::FrameTooShort(FRAME_HEADER_LEN, buf.len())); } - let frame_type = buf.get_u8(); - let stream_id = buf.get_u64(); - let payload_len = buf.get_u32(); - let flags = buf.get_u8(); + // Read header fields without consuming (use slicing instead of get_u*) + let frame_type = buf[0]; + let stream_id = u64::from_be_bytes(buf[1..9].try_into().unwrap()); + let payload_len = u32::from_be_bytes(buf[9..13].try_into().unwrap()); + let flags = buf[13]; if payload_len > MAX_PAYLOAD_LEN { return Err(FrameError::PayloadOverflow(payload_len)); } - if buf.len() < payload_len as usize { - // Not enough data yet — put header back and wait - // We can't easily un-get, so reconstruct - let mut restored = BytesMut::with_capacity(FRAME_HEADER_LEN + buf.len()); - restored.put_u8(frame_type); - restored.put_u64(stream_id); - restored.put_u32(payload_len); - restored.put_u8(flags); - restored.put_slice(buf.split_to(buf.len()).as_ref()); - // Actually, since we consumed the header, we need to return the error - // and the caller will retry. We restore nothing — the caller reads again. - // This is a "need more data" condition. - return Err(FrameError::FrameTooShort( - FRAME_HEADER_LEN + payload_len as usize, - FRAME_HEADER_LEN + buf.len(), - )); + // Need the full payload — do NOT consume yet + let total_len = FRAME_HEADER_LEN + payload_len as usize; + if buf.len() < total_len { + return Err(FrameError::FrameTooShort(total_len, buf.len())); } - let payload = buf.split_to(payload_len as usize).freeze(); + // Only now, consume the entire frame at once + let raw = buf.split_to(total_len); Ok(( Frame { frame_type, stream_id, flags, - payload, + payload: raw[FRAME_HEADER_LEN..].to_vec().into(), }, payload_len as usize, )) @@ -461,6 +455,112 @@ mod tests { let mut buf = BytesMut::from(&[0u8; 10][..]); let result = Frame::deserialize(&mut buf); assert!(matches!(result, Err(FrameError::FrameTooShort(_, _)))); + // CRITICAL: Buffer must NOT be consumed on failure — bytes must remain + assert_eq!( + buf.len(), + 10, + "partial read must not consume bytes from buffer" + ); + } + + #[test] + fn frame_partial_payload_preserves_buffer() { + // Frame with 11-byte payload, feed only header + 5 payload bytes + let frame = Frame::data(1, b"01234567890"); + let serialized = frame.serialize(); + // Header (14) + partial payload (5) = 19 bytes + let partial = &serialized[..19]; + + let mut buf = BytesMut::from(partial); + let result = Frame::deserialize(&mut buf); + assert!(matches!(result, Err(FrameError::FrameTooShort(_, _)))); + // Buffer must still contain all 19 bytes — header must NOT be consumed + assert_eq!( + buf.len(), + 19, + "partial payload must not consume header bytes" + ); + } + + #[test] + fn frame_partial_then_complete() { + // Simulate TCP segmentation: header arrives, then partial payload, then rest + let frame = Frame::data(3, b"hello world"); + let serialized = frame.serialize(); + let _total_len = serialized.len(); // 14 + 11 = 25 + + // Step 1: Feed only the header (14 bytes) + let mut buf = BytesMut::from(&serialized[..FRAME_HEADER_LEN]); + let result = Frame::deserialize(&mut buf); + assert!(matches!(result, Err(FrameError::FrameTooShort(_, _)))); + assert_eq!( + buf.len(), + FRAME_HEADER_LEN, + "header must not be consumed yet" + ); + + // Step 2: Append partial payload (5 bytes) + buf.extend_from_slice(&serialized[FRAME_HEADER_LEN..FRAME_HEADER_LEN + 5]); + let result = Frame::deserialize(&mut buf); + assert!(matches!(result, Err(FrameError::FrameTooShort(_, _)))); + assert_eq!( + buf.len(), + FRAME_HEADER_LEN + 5, + "partial payload must not be consumed yet" + ); + + // Step 3: Append remaining payload (6 bytes) + buf.extend_from_slice(&serialized[FRAME_HEADER_LEN + 5..]); + let (decoded, _) = Frame::deserialize(&mut buf).unwrap(); + assert_eq!(decoded.frame_type, FRAME_DATA); + assert_eq!(decoded.stream_id, 3); + assert_eq!(&decoded.payload[..], b"hello world"); + assert!(buf.is_empty()); + } + + #[test] + fn frame_zero_payload_header_only() { + // CLOSE frame has 0-byte payload; header alone should succeed + let frame = Frame::close(7); + let serialized = frame.serialize(); + + let mut buf = BytesMut::from(&serialized[..]); + let (decoded, _) = Frame::deserialize(&mut buf).unwrap(); + assert_eq!(decoded.frame_type, FRAME_CLOSE); + assert_eq!(decoded.stream_id, 7); + assert!(decoded.payload.is_empty()); + } + + #[test] + fn frame_partial_header_then_complete() { + // Feed header in chunks, then payload + let frame = Frame::data(5, b"test"); + let serialized = frame.serialize(); + + // Step 1: Feed first 7 bytes of header + let mut buf = BytesMut::from(&serialized[..7]); + let result = Frame::deserialize(&mut buf); + assert!(matches!(result, Err(FrameError::FrameTooShort(_, _)))); + assert_eq!(buf.len(), 7, "partial header must not be consumed"); + + // Step 2: Feed rest of header + buf.extend_from_slice(&serialized[7..FRAME_HEADER_LEN]); + let result = Frame::deserialize(&mut buf); + // Still not enough — header is complete but payload may be missing + // In this case, header is complete and payload is 4 bytes, but none arrived yet + assert!(matches!(result, Err(FrameError::FrameTooShort(_, _)))); + assert_eq!( + buf.len(), + FRAME_HEADER_LEN, + "full header must not be consumed yet" + ); + + // Step 3: Feed payload + buf.extend_from_slice(&serialized[FRAME_HEADER_LEN..]); + let (decoded, _) = Frame::deserialize(&mut buf).unwrap(); + assert_eq!(decoded.frame_type, FRAME_DATA); + assert_eq!(decoded.stream_id, 5); + assert_eq!(&decoded.payload[..], b"test"); } #[test] @@ -555,4 +655,147 @@ mod tests { fn frame_header_size() { assert_eq!(FRAME_HEADER_LEN, 14); } + + /// Custom async reader that returns data chunk by chunk. + struct ChunkedReader { + chunks: Vec>, + chunk_idx: usize, + pos_in_chunk: usize, + } + + impl ChunkedReader { + fn new(chunks: Vec>) -> Self { + ChunkedReader { + chunks, + chunk_idx: 0, + pos_in_chunk: 0, + } + } + } + + impl tokio::io::AsyncRead for ChunkedReader { + fn poll_read( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + buf: &mut tokio::io::ReadBuf<'_>, + ) -> std::task::Poll> { + let this = self.get_mut(); + if this.chunk_idx >= this.chunks.len() { + return std::task::Poll::Ready(Ok(())); + } + let chunk = &this.chunks[this.chunk_idx]; + let available = &chunk[this.pos_in_chunk..]; + let to_copy = available.len().min(buf.remaining()); + buf.put_slice(&available[..to_copy]); + this.pos_in_chunk += to_copy; + if this.pos_in_chunk >= chunk.len() { + this.chunk_idx += 1; + this.pos_in_chunk = 0; + } + std::task::Poll::Ready(Ok(())) + } + } + + #[tokio::test] + async fn frame_reader_header_then_payload_chunks() { + // FrameReader receives header chunk, then payload chunk separately + let frame = Frame::data(2, b"chunked data"); + let serialized = frame.serialize(); + + let chunk1 = serialized[..FRAME_HEADER_LEN].to_vec(); + let chunk2 = serialized[FRAME_HEADER_LEN..].to_vec(); + + let mut reader = FrameReader::new(); + let mut cursor = ChunkedReader::new(vec![chunk1, chunk2]); + let result = reader.read_frame(&mut cursor).await.unwrap().unwrap(); + + assert_eq!(result.frame_type, FRAME_DATA); + assert_eq!(result.stream_id, 2); + assert_eq!(&result.payload[..], b"chunked data"); + } + + #[tokio::test] + async fn frame_reader_three_chunk_split() { + // Header only, partial payload, remaining payload + let frame = Frame::data(4, b"hello world"); + let serialized = frame.serialize(); + + let chunk1 = serialized[..FRAME_HEADER_LEN].to_vec(); + let chunk2 = serialized[FRAME_HEADER_LEN..FRAME_HEADER_LEN + 5].to_vec(); + let chunk3 = serialized[FRAME_HEADER_LEN + 5..].to_vec(); + + let mut reader = FrameReader::new(); + let mut cursor = ChunkedReader::new(vec![chunk1, chunk2, chunk3]); + let result = reader.read_frame(&mut cursor).await.unwrap().unwrap(); + + assert_eq!(result.frame_type, FRAME_DATA); + assert_eq!(result.stream_id, 4); + assert_eq!(&result.payload[..], b"hello world"); + } + + #[tokio::test] + async fn frame_reader_header_split_across_chunks() { + // Header itself is split across multiple reads + let frame = Frame::close(10); + let serialized = frame.serialize(); + + let chunk1 = serialized[..5].to_vec(); + let chunk2 = serialized[5..FRAME_HEADER_LEN].to_vec(); + + let mut reader = FrameReader::new(); + let mut cursor = ChunkedReader::new(vec![chunk1, chunk2]); + let result = reader.read_frame(&mut cursor).await.unwrap().unwrap(); + + assert_eq!(result.frame_type, FRAME_CLOSE); + assert_eq!(result.stream_id, 10); + assert!(result.payload.is_empty()); + } + + #[tokio::test] + async fn frame_reader_multiple_frames_chunked() { + // Two frames, each split into chunks + let f1 = Frame::data(1, b"A"); + let f2 = Frame::data(3, b"B"); + + let mut all = Vec::new(); + all.extend_from_slice(&f1.serialize()); + all.extend_from_slice(&f2.serialize()); + + // Split into tiny 7-byte chunks + let mut chunks = Vec::new(); + let mut i = 0; + while i < all.len() { + let end = std::cmp::min(i + 7, all.len()); + chunks.push(all[i..end].to_vec()); + i = end; + } + + let mut reader = FrameReader::new(); + let mut cursor = ChunkedReader::new(chunks); + + let r1 = reader.read_frame(&mut cursor).await.unwrap().unwrap(); + assert_eq!(r1.stream_id, 1); + assert_eq!(&r1.payload[..], b"A"); + + let r2 = reader.read_frame(&mut cursor).await.unwrap().unwrap(); + assert_eq!(r2.stream_id, 3); + assert_eq!(&r2.payload[..], b"B"); + } + + #[tokio::test] + async fn frame_reader_zero_payload_chunked() { + // CLOSE frame (zero payload) delivered as header-only in two chunks + let frame = Frame::close(6); + let serialized = frame.serialize(); + + let chunk1 = serialized[..7].to_vec(); + let chunk2 = serialized[7..].to_vec(); + + let mut reader = FrameReader::new(); + let mut cursor = ChunkedReader::new(vec![chunk1, chunk2]); + let result = reader.read_frame(&mut cursor).await.unwrap().unwrap(); + + assert_eq!(result.frame_type, FRAME_CLOSE); + assert_eq!(result.stream_id, 6); + } }