use bytes::{Buf, Bytes, BytesMut};
use std::io;
use std::io::Write;
use tokio::io::{AsyncWrite, AsyncWriteExt};
use tokio_util::codec::Decoder;
use velo_ext::MessageType;
const SCHEMA_VERSION_V1: u16 = 1;
pub(crate) const DEFAULT_MAX_FRAME_SIZE: u32 = 16 * 1024 * 1024;
pub(crate) const MIN_HEADER_SIZE: usize = 2 + 1 + 4 + 4;
pub(crate) const COALESCE_THRESHOLD: usize = 64 * 1024;
pub(crate) const DIRECT_PREFIX_CAP: usize = 256;
#[inline]
pub(crate) fn stage_direct_prefix(
preamble: &[u8; MIN_HEADER_SIZE],
header: &[u8],
prefix: &mut [u8; DIRECT_PREFIX_CAP],
) -> Option<usize> {
let len = MIN_HEADER_SIZE + header.len();
if len > DIRECT_PREFIX_CAP {
return None;
}
prefix[..MIN_HEADER_SIZE].copy_from_slice(preamble);
prefix[MIN_HEADER_SIZE..len].copy_from_slice(header);
Some(len)
}
pub(crate) const DEFAULT_SHRINK_THRESHOLD: usize = 8 * 1024 * 1024;
pub(crate) const SHRINK_RESET_CAPACITY: usize = 256 * 1024;
pub(crate) fn parse_shrink_threshold(raw: Option<&str>) -> usize {
raw.and_then(|s| s.parse::<usize>().ok())
.unwrap_or(DEFAULT_SHRINK_THRESHOLD)
}
#[inline]
pub(crate) fn maybe_shrink_read_buffer(
buf: &mut BytesMut,
threshold: usize,
last_frame_total_size: usize,
) {
if buf.is_empty()
&& buf.capacity() > threshold
&& buf.capacity() > 2 * SHRINK_RESET_CAPACITY
&& last_frame_total_size.saturating_mul(2) <= buf.capacity()
{
*buf = BytesMut::with_capacity(SHRINK_RESET_CAPACITY);
}
}
#[derive(Debug, Clone)]
pub struct TcpFrameCodec {
state: DecodeState,
max_frame_size: u32,
}
#[derive(Debug, Clone, Copy)]
enum DecodeState {
AwaitingHeader,
AwaitingData {
frame_type: MessageType,
header_len: u32,
payload_len: u32,
},
}
impl TcpFrameCodec {
pub fn new() -> Self {
Self {
state: DecodeState::AwaitingHeader,
max_frame_size: DEFAULT_MAX_FRAME_SIZE,
}
}
pub fn with_max_frame_size(max_frame_size: u32) -> Self {
Self {
state: DecodeState::AwaitingHeader,
max_frame_size,
}
}
#[inline]
pub fn build_preamble(
msg_type: MessageType,
header_len: u32,
payload_len: u32,
) -> io::Result<[u8; MIN_HEADER_SIZE]> {
Self::validate_lengths(header_len, payload_len)?;
let mut preamble = [0u8; MIN_HEADER_SIZE];
preamble[0..2].copy_from_slice(&SCHEMA_VERSION_V1.to_be_bytes());
preamble[2] = msg_type.as_u8();
preamble[3..7].copy_from_slice(&header_len.to_be_bytes());
preamble[7..11].copy_from_slice(&payload_len.to_be_bytes());
Ok(preamble)
}
#[inline]
pub fn parse_message_type_from_preamble(preamble: &[u8]) -> io::Result<MessageType> {
if preamble.len() < MIN_HEADER_SIZE {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Preamble too short",
));
}
let schema_version = u16::from_be_bytes([preamble[0], preamble[1]]);
if schema_version != SCHEMA_VERSION_V1 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"Unsupported schema version: {} (expected {})",
schema_version, SCHEMA_VERSION_V1
),
));
}
MessageType::from_u8(preamble[2]).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("Invalid message type: {}", preamble[2]),
)
})
}
#[inline]
pub async fn encode_frame<W: AsyncWrite + Unpin>(
writer: &mut W,
msg_type: MessageType,
header: &[u8],
payload: &[u8],
) -> tokio::io::Result<()> {
let preamble = Self::build_preamble(msg_type, header.len() as u32, payload.len() as u32)?;
if header.len() + payload.len() <= COALESCE_THRESHOLD {
let mut buf = BytesMut::with_capacity(MIN_HEADER_SIZE + header.len() + payload.len());
buf.extend_from_slice(&preamble);
buf.extend_from_slice(header);
buf.extend_from_slice(payload);
writer.write_all(&buf).await?;
} else {
let mut prefix = [0u8; DIRECT_PREFIX_CAP];
match stage_direct_prefix(&preamble, header, &mut prefix) {
Some(len) => {
writer.write_all(&prefix[..len]).await?;
writer.write_all(payload).await?;
}
None => {
writer.write_all(&preamble).await?;
writer.write_all(header).await?;
writer.write_all(payload).await?;
}
}
}
Ok(())
}
#[inline]
pub(crate) fn append_frame(
buf: &mut BytesMut,
msg_type: MessageType,
header: &[u8],
payload: &[u8],
) -> io::Result<()> {
let preamble = Self::build_preamble(msg_type, header.len() as u32, payload.len() as u32)?;
buf.reserve(MIN_HEADER_SIZE + header.len() + payload.len());
buf.extend_from_slice(&preamble);
buf.extend_from_slice(header);
buf.extend_from_slice(payload);
Ok(())
}
#[inline]
pub fn encode_frame_sync<W: Write>(
writer: &mut W,
msg_type: MessageType,
header: &[u8],
payload: &[u8],
) -> std::io::Result<()> {
let preamble = Self::build_preamble(msg_type, header.len() as u32, payload.len() as u32)?;
if header.len() + payload.len() <= COALESCE_THRESHOLD {
let mut buf = BytesMut::with_capacity(MIN_HEADER_SIZE + header.len() + payload.len());
buf.extend_from_slice(&preamble);
buf.extend_from_slice(header);
buf.extend_from_slice(payload);
writer.write_all(&buf)?;
} else {
let mut prefix = [0u8; DIRECT_PREFIX_CAP];
match stage_direct_prefix(&preamble, header, &mut prefix) {
Some(len) => {
writer.write_all(&prefix[..len])?;
writer.write_all(payload)?;
}
None => {
writer.write_all(&preamble)?;
writer.write_all(header)?;
writer.write_all(payload)?;
}
}
}
Ok(())
}
fn validate_lengths(header_len: u32, payload_len: u32) -> io::Result<()> {
Self::validate_lengths_limit(header_len, payload_len, DEFAULT_MAX_FRAME_SIZE)
}
fn validate_lengths_limit(
header_len: u32,
payload_len: u32,
max_frame_size: u32,
) -> io::Result<()> {
let total_len = header_len
.checked_add(payload_len)
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "Frame size overflow"))?;
if total_len > max_frame_size {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"Frame size {} exceeds maximum {}",
total_len, max_frame_size
),
));
}
Ok(())
}
}
impl Default for TcpFrameCodec {
fn default() -> Self {
Self::new()
}
}
impl Decoder for TcpFrameCodec {
type Item = (MessageType, Bytes, Bytes);
type Error = io::Error;
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
loop {
match self.state {
DecodeState::AwaitingHeader => {
if src.len() < MIN_HEADER_SIZE {
return Ok(None);
}
let schema_version = u16::from_be_bytes([src[0], src[1]]);
let frame_type_byte = src[2];
let header_len = u32::from_be_bytes([src[3], src[4], src[5], src[6]]);
let payload_len = u32::from_be_bytes([src[7], src[8], src[9], src[10]]);
if schema_version != SCHEMA_VERSION_V1 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"Unsupported schema version: {} (expected {})",
schema_version, SCHEMA_VERSION_V1
),
));
}
let frame_type = MessageType::from_u8(frame_type_byte).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("Invalid frame type: {}", frame_type_byte),
)
})?;
Self::validate_lengths_limit(header_len, payload_len, self.max_frame_size)?;
src.advance(MIN_HEADER_SIZE);
self.state = DecodeState::AwaitingData {
frame_type,
header_len,
payload_len,
};
}
DecodeState::AwaitingData {
frame_type,
header_len,
payload_len,
..
} => {
let total_data_len = (header_len + payload_len) as usize;
if src.len() < total_data_len {
src.reserve(total_data_len - src.len());
return Ok(None);
}
let header = src.split_to(header_len as usize).freeze();
let payload = src.split_to(payload_len as usize).freeze();
self.state = DecodeState::AwaitingHeader;
return Ok(Some((frame_type, header, payload)));
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio_util::codec::Framed;
async fn encode_frame_to_bytes(
msg_type: MessageType,
header: &[u8],
payload: &[u8],
) -> io::Result<Vec<u8>> {
let mut buf = Vec::new();
TcpFrameCodec::encode_frame(&mut buf, msg_type, header, payload).await?;
Ok(buf)
}
fn encode_frame_to_bytes_sync(
msg_type: MessageType,
header: &[u8],
payload: &[u8],
) -> io::Result<Vec<u8>> {
let mut buf = Vec::new();
TcpFrameCodec::encode_frame_sync(&mut buf, msg_type, header, payload)?;
Ok(buf)
}
fn create_unsafe_frame(
schema_version: u16,
frame_type: MessageType,
header: &[u8],
payload: &[u8],
) -> BytesMut {
let mut buf = BytesMut::new();
buf.extend_from_slice(&schema_version.to_be_bytes());
buf.extend_from_slice(&[frame_type.as_u8()]);
buf.extend_from_slice(&(header.len() as u32).to_be_bytes());
buf.extend_from_slice(&(payload.len() as u32).to_be_bytes());
buf.extend_from_slice(header);
buf.extend_from_slice(payload);
buf
}
#[test]
fn test_decode_message_frame() {
let mut codec = TcpFrameCodec::new();
let header = b"test-header";
let payload = b"test-payload-data";
let framed = encode_frame_to_bytes_sync(MessageType::Message, header, payload).unwrap();
let mut buf = BytesMut::from(&framed[..]);
let result = codec.decode(&mut buf).unwrap();
assert!(result.is_some());
let (msg_type, decoded_header, decoded_payload) = result.unwrap();
assert_eq!(msg_type, MessageType::Message);
assert_eq!(decoded_header, Bytes::from(header.as_ref()));
assert_eq!(decoded_payload, Bytes::from(payload.as_ref()));
}
#[test]
fn test_decode_all_frame_types() {
let frame_types = [
MessageType::Message,
MessageType::Response,
MessageType::Ack,
MessageType::Event,
];
for frame_type in &frame_types {
let mut codec = TcpFrameCodec::new();
let header = b"header";
let payload = b"payload";
let framed = encode_frame_to_bytes_sync(*frame_type, header, payload).unwrap();
let mut buf = BytesMut::from(&framed[..]);
let result = codec.decode(&mut buf).unwrap();
assert!(result.is_some());
let (decoded_type, _, _) = result.unwrap();
assert_eq!(decoded_type, *frame_type);
}
}
#[test]
fn test_decode_empty_payload() {
let mut codec = TcpFrameCodec::new();
let header = b"ack-header";
let payload = b"";
let framed = encode_frame_to_bytes_sync(MessageType::Ack, header, payload).unwrap();
let mut buf = BytesMut::from(&framed[..]);
let result = codec.decode(&mut buf).unwrap();
assert!(result.is_some());
let (msg_type, decoded_header, decoded_payload) = result.unwrap();
assert_eq!(msg_type, MessageType::Ack);
assert_eq!(&decoded_header[..], header);
assert_eq!(decoded_payload.len(), 0);
}
#[test]
fn test_decode_partial_frame() {
let mut codec = TcpFrameCodec::new();
let header = b"test-header";
let payload = b"test-payload";
let full_frame = encode_frame_to_bytes_sync(MessageType::Message, header, payload).unwrap();
let mut buf = BytesMut::from(&full_frame[..5]);
let result = codec.decode(&mut buf).unwrap();
assert!(result.is_none());
buf.extend_from_slice(&full_frame[5..MIN_HEADER_SIZE]);
let result = codec.decode(&mut buf).unwrap();
assert!(result.is_none());
buf.extend_from_slice(&full_frame[MIN_HEADER_SIZE..]);
let result = codec.decode(&mut buf).unwrap();
assert!(result.is_some());
let (msg_type, decoded_header, decoded_payload) = result.unwrap();
assert_eq!(msg_type, MessageType::Message);
assert_eq!(&decoded_header[..], header);
assert_eq!(&decoded_payload[..], payload);
}
#[test]
fn test_decode_invalid_schema_version() {
let mut codec = TcpFrameCodec::new();
let header = b"header";
let payload = b"payload";
let mut buf = create_unsafe_frame(999, MessageType::Message, header, payload);
let result = codec.decode(&mut buf);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("Unsupported schema version")
);
}
#[test]
fn test_decode_invalid_frame_type() {
let mut codec = TcpFrameCodec::new();
let mut buf = BytesMut::new();
buf.extend_from_slice(&SCHEMA_VERSION_V1.to_be_bytes());
buf.extend_from_slice(&[255u8]); buf.extend_from_slice(&10u32.to_be_bytes()); buf.extend_from_slice(&10u32.to_be_bytes());
let result = codec.decode(&mut buf);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("Invalid frame type")
);
}
#[test]
fn test_decode_frame_too_large() {
let mut codec = TcpFrameCodec::new();
let mut buf = BytesMut::new();
buf.extend_from_slice(&SCHEMA_VERSION_V1.to_be_bytes());
buf.extend_from_slice(&[MessageType::Message.as_u8()]);
buf.extend_from_slice(&(DEFAULT_MAX_FRAME_SIZE / 2 + 1).to_be_bytes());
buf.extend_from_slice(&(DEFAULT_MAX_FRAME_SIZE / 2 + 1).to_be_bytes());
let result = codec.decode(&mut buf);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("exceeds maximum"));
}
#[test]
fn test_decode_multiple_frames() {
let mut codec = TcpFrameCodec::new();
let mut buf = BytesMut::new();
let frame1 =
encode_frame_to_bytes_sync(MessageType::Message, b"header1", b"payload1").unwrap();
let frame2 =
encode_frame_to_bytes_sync(MessageType::Response, b"header2", b"payload2").unwrap();
buf.extend_from_slice(&frame1);
buf.extend_from_slice(&frame2);
let result = codec.decode(&mut buf).unwrap();
assert!(result.is_some());
let (msg_type, header, payload) = result.unwrap();
assert_eq!(msg_type, MessageType::Message);
assert_eq!(&header[..], b"header1");
assert_eq!(&payload[..], b"payload1");
let result = codec.decode(&mut buf).unwrap();
assert!(result.is_some());
let (msg_type, header, payload) = result.unwrap();
assert_eq!(msg_type, MessageType::Response);
assert_eq!(&header[..], b"header2");
assert_eq!(&payload[..], b"payload2");
assert!(buf.is_empty());
}
#[test]
fn test_zero_copy_bytes_share_buffer() {
let mut codec = TcpFrameCodec::new();
let header = b"shared-header";
let payload = b"shared-payload";
let framed = encode_frame_to_bytes_sync(MessageType::Message, header, payload).unwrap();
let mut buf = BytesMut::from(&framed[..]);
let result = codec.decode(&mut buf).unwrap().unwrap();
let (_, decoded_header, decoded_payload) = result;
assert_eq!(&decoded_header[..], header);
assert_eq!(&decoded_payload[..], payload);
let header_clone = decoded_header.clone();
let payload_clone = decoded_payload.clone();
assert_eq!(decoded_header, header_clone);
assert_eq!(decoded_payload, payload_clone);
}
#[test]
fn test_encode_frame() {
let header = b"test-header";
let payload = b"test-payload";
let framed = encode_frame_to_bytes_sync(MessageType::Message, header, payload).unwrap();
assert_eq!(framed.len(), MIN_HEADER_SIZE + header.len() + payload.len());
assert_eq!(
u16::from_be_bytes([framed[0], framed[1]]),
SCHEMA_VERSION_V1
);
assert_eq!(framed[2], MessageType::Message.as_u8());
assert_eq!(
u32::from_be_bytes([framed[3], framed[4], framed[5], framed[6]]),
header.len() as u32
);
assert_eq!(
u32::from_be_bytes([framed[7], framed[8], framed[9], framed[10]]),
payload.len() as u32
);
assert_eq!(
&framed[MIN_HEADER_SIZE..MIN_HEADER_SIZE + header.len()],
header
);
assert_eq!(&framed[MIN_HEADER_SIZE + header.len()..], payload);
}
#[test]
fn test_encode_all_message_types() {
let header = b"header";
let payload = b"payload";
for msg_type in &[
MessageType::Message,
MessageType::Response,
MessageType::Ack,
MessageType::Event,
] {
let framed = encode_frame_to_bytes_sync(*msg_type, header, payload).unwrap();
assert_eq!(framed[2], msg_type.as_u8());
}
}
#[test]
fn test_encode_empty_payload() {
let header = b"ack-header";
let payload = b"";
let framed = encode_frame_to_bytes_sync(MessageType::Ack, header, payload).unwrap();
assert_eq!(framed.len(), MIN_HEADER_SIZE + header.len());
assert_eq!(
u32::from_be_bytes([framed[7], framed[8], framed[9], framed[10]]),
0
);
}
#[test]
fn test_encode_frame_too_large() {
let header = vec![0u8; (DEFAULT_MAX_FRAME_SIZE / 2 + 1) as usize];
let payload = vec![0u8; (DEFAULT_MAX_FRAME_SIZE / 2 + 1) as usize];
let result = encode_frame_to_bytes_sync(MessageType::Message, &header, &payload);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("exceeds maximum"));
}
#[test]
fn test_round_trip_encode_decode() {
let mut codec = TcpFrameCodec::new();
let header = b"round-trip-header";
let payload = b"round-trip-payload-data";
let framed = encode_frame_to_bytes_sync(MessageType::Response, header, payload).unwrap();
let mut buf = BytesMut::from(&framed[..]);
let result = codec.decode(&mut buf).unwrap();
assert!(result.is_some());
let (msg_type, decoded_header, decoded_payload) = result.unwrap();
assert_eq!(msg_type, MessageType::Response);
assert_eq!(&decoded_header[..], header);
assert_eq!(&decoded_payload[..], payload);
}
#[test]
fn test_round_trip_all_types() {
let types = [
MessageType::Message,
MessageType::Response,
MessageType::Ack,
MessageType::Event,
];
for msg_type in &types {
let mut codec = TcpFrameCodec::new();
let header = b"header";
let payload = b"payload";
let framed = encode_frame_to_bytes_sync(*msg_type, header, payload).unwrap();
let mut buf = BytesMut::from(&framed[..]);
let result = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(result.0, *msg_type);
assert_eq!(&result.1[..], header);
assert_eq!(&result.2[..], payload);
}
}
#[test]
fn test_encode_frame_sync() {
let header = b"sync-header";
let payload = b"sync-payload";
let framed = encode_frame_to_bytes_sync(MessageType::Message, header, payload).unwrap();
assert_eq!(framed.len(), MIN_HEADER_SIZE + header.len() + payload.len());
assert_eq!(
u16::from_be_bytes([framed[0], framed[1]]),
SCHEMA_VERSION_V1
);
assert_eq!(framed[2], MessageType::Message.as_u8());
assert_eq!(
u32::from_be_bytes([framed[3], framed[4], framed[5], framed[6]]),
header.len() as u32
);
assert_eq!(
u32::from_be_bytes([framed[7], framed[8], framed[9], framed[10]]),
payload.len() as u32
);
assert_eq!(
&framed[MIN_HEADER_SIZE..MIN_HEADER_SIZE + header.len()],
header
);
assert_eq!(&framed[MIN_HEADER_SIZE + header.len()..], payload);
}
#[test]
fn test_sync_async_produce_same_output() {
let header = b"test-header";
let payload = b"test-payload";
let sync_framed =
encode_frame_to_bytes_sync(MessageType::Response, header, payload).unwrap();
let async_framed = tokio::runtime::Runtime::new()
.unwrap()
.block_on(encode_frame_to_bytes(
MessageType::Response,
header,
payload,
))
.unwrap();
assert_eq!(sync_framed, async_framed);
}
fn framed_with_grown_buffer(
capacity_bytes: usize,
) -> Framed<tokio::io::DuplexStream, TcpFrameCodec> {
let (read_half, _write_half) = tokio::io::duplex(64);
let mut framed = Framed::new(read_half, TcpFrameCodec::new());
framed.read_buffer_mut().resize(capacity_bytes, 0);
framed.read_buffer_mut().clear(); framed
}
#[tokio::test]
async fn test_maybe_shrink_resets_oversized_empty_buffer() {
let mut framed = framed_with_grown_buffer(2 * 1024 * 1024);
let pre_capacity = framed.read_buffer_mut().capacity();
assert!(pre_capacity > 2 * SHRINK_RESET_CAPACITY);
maybe_shrink_read_buffer(framed.read_buffer_mut(), 1024 * 1024, 1024);
let post_capacity = framed.read_buffer_mut().capacity();
assert!(
post_capacity < pre_capacity,
"expected capacity to drop: pre={} post={}",
pre_capacity,
post_capacity
);
assert_eq!(post_capacity, SHRINK_RESET_CAPACITY);
}
#[tokio::test]
async fn test_maybe_shrink_skips_when_not_empty() {
let (read_half, _write_half) = tokio::io::duplex(64);
let mut framed = Framed::new(read_half, TcpFrameCodec::new());
framed.read_buffer_mut().resize(2 * 1024 * 1024, 0xAB);
let pre_capacity = framed.read_buffer_mut().capacity();
maybe_shrink_read_buffer(framed.read_buffer_mut(), 1024 * 1024, 1024);
assert_eq!(framed.read_buffer_mut().capacity(), pre_capacity);
}
#[tokio::test]
async fn test_maybe_shrink_skips_under_threshold() {
let (read_half, _write_half) = tokio::io::duplex(64);
let mut framed = Framed::new(read_half, TcpFrameCodec::new());
let pre_capacity = framed.read_buffer_mut().capacity();
maybe_shrink_read_buffer(framed.read_buffer_mut(), DEFAULT_SHRINK_THRESHOLD, 1024);
assert_eq!(framed.read_buffer_mut().capacity(), pre_capacity);
}
#[tokio::test]
async fn test_maybe_shrink_skips_under_sustained_large_frames() {
let cap = 2 * 1024 * 1024;
let mut framed = framed_with_grown_buffer(cap);
let pre = framed.read_buffer_mut().capacity();
maybe_shrink_read_buffer(framed.read_buffer_mut(), 1024 * 1024, cap * 3 / 4);
assert_eq!(framed.read_buffer_mut().capacity(), pre);
}
#[tokio::test]
async fn test_maybe_shrink_skips_when_capacity_not_meaningfully_oversized() {
let cap = SHRINK_RESET_CAPACITY + 4096;
let mut framed = framed_with_grown_buffer(cap);
let pre = framed.read_buffer_mut().capacity();
maybe_shrink_read_buffer(framed.read_buffer_mut(), 1024, 100); assert_eq!(framed.read_buffer_mut().capacity(), pre);
}
#[test]
fn test_parse_shrink_threshold() {
assert_eq!(parse_shrink_threshold(Some("12345")), 12345);
assert_eq!(parse_shrink_threshold(Some("0")), 0);
assert_eq!(
parse_shrink_threshold(Some("not-a-number")),
DEFAULT_SHRINK_THRESHOLD
);
assert_eq!(parse_shrink_threshold(Some("")), DEFAULT_SHRINK_THRESHOLD);
assert_eq!(parse_shrink_threshold(None), DEFAULT_SHRINK_THRESHOLD);
}
#[test]
fn test_coalesce_and_segmented_paths_agree() {
let sizes = [
100usize,
COALESCE_THRESHOLD - 16,
COALESCE_THRESHOLD,
COALESCE_THRESHOLD + 1024,
];
for total in sizes {
let header_len = total / 2;
let payload_len = total - header_len;
let header: Vec<u8> = (0..header_len).map(|i| (i % 251) as u8).collect();
let payload: Vec<u8> = (0..payload_len).map(|i| (i % 253) as u8).collect();
let framed =
encode_frame_to_bytes_sync(MessageType::Message, &header, &payload).unwrap();
assert_eq!(framed.len(), MIN_HEADER_SIZE + header_len + payload_len);
let mut codec = TcpFrameCodec::new();
let mut buf = BytesMut::from(&framed[..]);
let (msg_type, decoded_header, decoded_payload) =
codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(msg_type, MessageType::Message);
assert_eq!(&decoded_header[..], header.as_slice());
assert_eq!(&decoded_payload[..], payload.as_slice());
}
}
#[derive(Default)]
struct CountingWriter {
data: Vec<u8>,
writes: usize,
}
impl Write for CountingWriter {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.writes += 1;
self.data.extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[test]
fn test_direct_path_stages_prefix_into_two_writes() {
let header = vec![0x11u8; 64];
let payload = vec![0xABu8; COALESCE_THRESHOLD + 1];
let mut w = CountingWriter::default();
TcpFrameCodec::encode_frame_sync(&mut w, MessageType::Message, &header, &payload).unwrap();
assert_eq!(w.writes, 2, "staged prefix + payload");
let reference =
create_unsafe_frame(SCHEMA_VERSION_V1, MessageType::Message, &header, &payload);
assert_eq!(w.data, reference.to_vec(), "wire bytes must be unchanged");
}
#[test]
fn test_direct_path_oversized_header_falls_back() {
let header = vec![0x22u8; DIRECT_PREFIX_CAP - MIN_HEADER_SIZE + 1];
let payload = vec![0xABu8; COALESCE_THRESHOLD + 1];
let mut w = CountingWriter::default();
TcpFrameCodec::encode_frame_sync(&mut w, MessageType::Message, &header, &payload).unwrap();
assert_eq!(
w.writes, 3,
"preamble, header, payload each their own write"
);
let reference =
create_unsafe_frame(SCHEMA_VERSION_V1, MessageType::Message, &header, &payload);
assert_eq!(w.data, reference.to_vec(), "wire bytes must be unchanged");
}
#[test]
fn test_decode_reserves_capacity_for_announced_frame() {
let mut codec = TcpFrameCodec::new();
let header_len = 64u32;
let payload_len = 4 * 1024 * 1024u32;
let mut buf = BytesMut::with_capacity(8 * 1024);
buf.extend_from_slice(&SCHEMA_VERSION_V1.to_be_bytes());
buf.extend_from_slice(&[MessageType::Message.as_u8()]);
buf.extend_from_slice(&header_len.to_be_bytes());
buf.extend_from_slice(&payload_len.to_be_bytes());
buf.extend_from_slice(&[0u8; 1000]);
assert!(codec.decode(&mut buf).unwrap().is_none());
let total = (header_len + payload_len) as usize;
assert!(
buf.capacity() >= total,
"decoder must reserve for the announced frame: capacity={} < {}",
buf.capacity(),
total
);
}
}