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:
+267
-24
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user