use bytes::{Buf, BufMut, Bytes, BytesMut};
use thiserror::Error;
pub const PROTOCOL_VERSION: u8 = 1;
pub const MAX_MESSAGE_SIZE: u32 = 16 * 1024 * 1024;
pub const FRAME_HEADER_SIZE: usize = 5;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FrameHeader {
pub version: u8,
pub length: u32,
}
impl FrameHeader {
pub fn new(length: u32) -> Self {
Self {
version: PROTOCOL_VERSION,
length,
}
}
pub fn encode(&self, dst: &mut BytesMut) {
dst.put_u8(self.version);
dst.put_u32(self.length);
}
pub fn decode(src: &[u8]) -> Option<Self> {
if src.len() < FRAME_HEADER_SIZE {
return None;
}
Some(Self {
version: src[0],
length: u32::from_be_bytes([src[1], src[2], src[3], src[4]]),
})
}
pub fn frame_size(&self) -> usize {
FRAME_HEADER_SIZE + self.length as usize
}
}
#[derive(Debug, Clone, Error, PartialEq, Eq)]
pub enum FrameError {
#[error("message too large: {0} bytes (max: {1})")]
MessageTooLarge(u32, u32),
#[error("invalid protocol version: {0}")]
InvalidVersion(u8),
#[error("incomplete frame")]
Incomplete,
#[error("empty message")]
EmptyMessage,
}
#[derive(Debug)]
pub struct FrameCodec {
max_size: u32,
read_buffer: BytesMut,
}
impl Default for FrameCodec {
fn default() -> Self {
Self::new()
}
}
impl FrameCodec {
pub fn new() -> Self {
Self {
max_size: MAX_MESSAGE_SIZE,
read_buffer: BytesMut::with_capacity(8192),
}
}
pub fn with_max_size(max_size: u32) -> Self {
Self {
max_size,
read_buffer: BytesMut::with_capacity(8192),
}
}
pub fn max_size(&self) -> u32 {
self.max_size
}
pub fn feed(&mut self, data: &[u8]) {
self.read_buffer.extend_from_slice(data);
}
pub fn has_complete_frame(&self) -> bool {
if self.read_buffer.len() < FRAME_HEADER_SIZE {
return false;
}
if let Some(header) = FrameHeader::decode(&self.read_buffer) {
if header.version != PROTOCOL_VERSION || header.length > self.max_size {
return true;
}
self.read_buffer.len() >= header.frame_size()
} else {
false
}
}
pub fn buffer_size(&self) -> usize {
self.read_buffer.len()
}
pub fn clear(&mut self) {
self.read_buffer.clear();
}
pub fn decode(&mut self) -> Result<Option<Bytes>, FrameError> {
let mut probe = self.read_buffer.clone();
match self.decode_from(&mut probe) {
Ok(Some(bytes)) => {
let consumed = FRAME_HEADER_SIZE + bytes.len();
self.read_buffer.advance(consumed);
Ok(Some(bytes))
}
Ok(None) => Ok(None),
Err(error) => {
self.read_buffer.clear();
Err(error)
}
}
}
pub fn decode_from(&self, src: &mut BytesMut) -> Result<Option<Bytes>, FrameError> {
if src.len() < FRAME_HEADER_SIZE {
return Ok(None);
}
let header = FrameHeader::decode(src).ok_or(FrameError::Incomplete)?;
if header.version != PROTOCOL_VERSION {
return Err(FrameError::InvalidVersion(header.version));
}
if header.length > self.max_size {
return Err(FrameError::MessageTooLarge(header.length, self.max_size));
}
let frame_size = header.frame_size();
if src.len() < frame_size {
return Ok(None);
}
src.advance(FRAME_HEADER_SIZE);
let payload = src.split_to(header.length as usize).freeze();
Ok(Some(payload))
}
pub fn encode(&self, msg: &[u8], dst: &mut BytesMut) -> Result<(), FrameError> {
let len = msg.len() as u32;
if len > self.max_size {
return Err(FrameError::MessageTooLarge(len, self.max_size));
}
dst.reserve(FRAME_HEADER_SIZE + msg.len());
let header = FrameHeader::new(len);
header.encode(dst);
dst.extend_from_slice(msg);
Ok(())
}
pub fn encode_to_bytes(&self, msg: &[u8]) -> Result<Bytes, FrameError> {
let mut dst = BytesMut::with_capacity(FRAME_HEADER_SIZE + msg.len());
self.encode(msg, &mut dst)?;
Ok(dst.freeze())
}
}
pub fn frame_message(msg: &[u8]) -> Result<Bytes, FrameError> {
FrameCodec::new().encode_to_bytes(msg)
}
pub fn unframe_message(data: &[u8]) -> Result<Option<Bytes>, FrameError> {
let mut src = BytesMut::from(data);
FrameCodec::new().decode_from(&mut src)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_frame_header_encode_decode() {
let header = FrameHeader::new(1234);
let mut buf = BytesMut::new();
header.encode(&mut buf);
assert_eq!(buf.len(), FRAME_HEADER_SIZE);
let decoded = FrameHeader::decode(&buf).unwrap();
assert_eq!(decoded.version, PROTOCOL_VERSION);
assert_eq!(decoded.length, 1234);
}
#[test]
fn test_frame_header_frame_size() {
let header = FrameHeader::new(100);
assert_eq!(header.frame_size(), FRAME_HEADER_SIZE + 100);
}
#[test]
fn test_frame_codec_encode_decode() {
let codec = FrameCodec::new();
let message = b"Hello, DCP!";
let mut encoded = BytesMut::new();
codec.encode(message, &mut encoded).unwrap();
let decoded = codec.decode_from(&mut encoded).unwrap().unwrap();
assert_eq!(&decoded[..], message);
}
#[test]
fn test_frame_codec_round_trip() {
let codec = FrameCodec::new();
let messages = vec![
b"".to_vec(),
b"short".to_vec(),
b"a longer message with more content".to_vec(),
vec![0u8; 1000], vec![0xFF; 10000], ];
for msg in messages {
let framed = codec.encode_to_bytes(&msg).unwrap();
let mut src = BytesMut::from(&framed[..]);
let decoded = codec.decode_from(&mut src).unwrap().unwrap();
assert_eq!(&decoded[..], &msg[..]);
}
}
#[test]
fn test_frame_codec_partial_message() {
let codec = FrameCodec::new();
let message = b"Complete message";
let framed = codec.encode_to_bytes(message).unwrap();
let mut partial = BytesMut::from(&framed[..FRAME_HEADER_SIZE]);
let result = codec.decode_from(&mut partial).unwrap();
assert!(result.is_none());
let mut partial = BytesMut::from(&framed[..FRAME_HEADER_SIZE + 5]);
let result = codec.decode_from(&mut partial).unwrap();
assert!(result.is_none());
let mut full = BytesMut::from(&framed[..]);
let result = codec.decode_from(&mut full).unwrap();
assert!(result.is_some());
assert_eq!(&result.unwrap()[..], message);
}
#[test]
fn test_frame_codec_message_too_large() {
let codec = FrameCodec::with_max_size(100);
let message = vec![0u8; 200];
let result = codec.encode_to_bytes(&message);
assert!(matches!(result, Err(FrameError::MessageTooLarge(200, 100))));
}
#[test]
fn test_frame_codec_invalid_version() {
let mut data = BytesMut::new();
data.put_u8(99); data.put_u32(5);
data.extend_from_slice(b"hello");
let codec = FrameCodec::new();
let result = codec.decode_from(&mut data);
assert!(matches!(result, Err(FrameError::InvalidVersion(99))));
}
#[test]
fn test_frame_codec_feed_and_decode() {
let mut codec = FrameCodec::new();
let message = b"Test message";
let framed = codec.encode_to_bytes(message).unwrap();
codec.feed(&framed[..3]); assert!(!codec.has_complete_frame());
codec.feed(&framed[3..FRAME_HEADER_SIZE]); assert!(!codec.has_complete_frame());
codec.feed(&framed[FRAME_HEADER_SIZE..]); assert!(codec.has_complete_frame());
let decoded = codec.decode().unwrap().unwrap();
assert_eq!(&decoded[..], message);
assert_eq!(codec.buffer_size(), 0);
}
#[test]
fn test_frame_codec_multiple_messages() {
let mut codec = FrameCodec::new();
let msg1 = b"First";
let msg2 = b"Second";
let framed1 = codec.encode_to_bytes(msg1).unwrap();
let framed2 = codec.encode_to_bytes(msg2).unwrap();
codec.feed(&framed1);
codec.feed(&framed2);
let decoded1 = codec.decode().unwrap().unwrap();
assert_eq!(&decoded1[..], msg1);
let decoded2 = codec.decode().unwrap().unwrap();
assert_eq!(&decoded2[..], msg2);
assert!(!codec.has_complete_frame());
}
#[test]
fn test_convenience_functions() {
let message = b"Quick test";
let framed = frame_message(message).unwrap();
let unframed = unframe_message(&framed).unwrap().unwrap();
assert_eq!(&unframed[..], message);
}
#[test]
fn test_frame_header_decode_insufficient_data() {
let data = [0u8; 3]; assert!(FrameHeader::decode(&data).is_none());
}
#[test]
fn test_max_message_size_validation() {
let mut data = BytesMut::new();
data.put_u8(PROTOCOL_VERSION);
data.put_u32(MAX_MESSAGE_SIZE + 1);
let codec = FrameCodec::new();
let result = codec.decode_from(&mut data);
assert!(matches!(result, Err(FrameError::MessageTooLarge(_, _))));
}
}