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.
|
/// 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> {
|
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 {
|
if buf.len() < FRAME_HEADER_LEN {
|
||||||
return Err(FrameError::FrameTooShort(FRAME_HEADER_LEN, buf.len()));
|
return Err(FrameError::FrameTooShort(FRAME_HEADER_LEN, buf.len()));
|
||||||
}
|
}
|
||||||
|
|
||||||
let frame_type = buf.get_u8();
|
// Read header fields without consuming (use slicing instead of get_u*)
|
||||||
let stream_id = buf.get_u64();
|
let frame_type = buf[0];
|
||||||
let payload_len = buf.get_u32();
|
let stream_id = u64::from_be_bytes(buf[1..9].try_into().unwrap());
|
||||||
let flags = buf.get_u8();
|
let payload_len = u32::from_be_bytes(buf[9..13].try_into().unwrap());
|
||||||
|
let flags = buf[13];
|
||||||
|
|
||||||
if payload_len > MAX_PAYLOAD_LEN {
|
if payload_len > MAX_PAYLOAD_LEN {
|
||||||
return Err(FrameError::PayloadOverflow(payload_len));
|
return Err(FrameError::PayloadOverflow(payload_len));
|
||||||
}
|
}
|
||||||
|
|
||||||
if buf.len() < payload_len as usize {
|
// Need the full payload — do NOT consume yet
|
||||||
// Not enough data yet — put header back and wait
|
let total_len = FRAME_HEADER_LEN + payload_len as usize;
|
||||||
// We can't easily un-get, so reconstruct
|
if buf.len() < total_len {
|
||||||
let mut restored = BytesMut::with_capacity(FRAME_HEADER_LEN + buf.len());
|
return Err(FrameError::FrameTooShort(total_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(),
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
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((
|
Ok((
|
||||||
Frame {
|
Frame {
|
||||||
frame_type,
|
frame_type,
|
||||||
stream_id,
|
stream_id,
|
||||||
flags,
|
flags,
|
||||||
payload,
|
payload: raw[FRAME_HEADER_LEN..].to_vec().into(),
|
||||||
},
|
},
|
||||||
payload_len as usize,
|
payload_len as usize,
|
||||||
))
|
))
|
||||||
@@ -461,6 +455,112 @@ mod tests {
|
|||||||
let mut buf = BytesMut::from(&[0u8; 10][..]);
|
let mut buf = BytesMut::from(&[0u8; 10][..]);
|
||||||
let result = Frame::deserialize(&mut buf);
|
let result = Frame::deserialize(&mut buf);
|
||||||
assert!(matches!(result, Err(FrameError::FrameTooShort(_, _))));
|
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]
|
#[test]
|
||||||
@@ -555,4 +655,147 @@ mod tests {
|
|||||||
fn frame_header_size() {
|
fn frame_header_size() {
|
||||||
assert_eq!(FRAME_HEADER_LEN, 14);
|
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