fix: prevent frame header corruption on partial reads

Frame::deserialize consumed header bytes from the buffer via get_u8/
get_u64/get_u32 even when the payload was incomplete, causing those
bytes to be permanently lost on retry. Replaced with slice-based
reading that never consumes until the complete frame (header + payload)
is available.

Added 9 new tests covering header-only, partial payload, remaining
payload, header-split, and multi-frame chunked scenarios.
Includes async FrameReader tests with chunked reads.
This commit is contained in:
c4ch3c4d3
2026-06-03 21:10:00 -06:00
parent d6a0a40696
commit b67248e14f
+267 -24
View File
@@ -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<Vec<u8>>,
chunk_idx: usize,
pos_in_chunk: usize,
}
impl ChunkedReader {
fn new(chunks: Vec<Vec<u8>>) -> 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<std::io::Result<()>> {
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);
}
}