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;
const DEFAULT_MAX_FRAME_SIZE: u32 = 16 * 1024 * 1024;
const MIN_HEADER_SIZE: usize = 2 + 1 + 4 + 4;
#[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)?;
writer.write_all(&preamble).await?;
writer.write_all(header).await?;
writer.write_all(payload).await?;
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)?;
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 {
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::*;
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);
}
}