use crate::error::ZmqError;
use crate::message::{Msg, MsgFlags};
use crate::protocol::zmtp::command::{ZMTP_FLAG_COMMAND, ZMTP_FLAG_LONG, ZMTP_FLAG_MORE};
use bytes::{Buf, BufMut, Bytes, BytesMut};
#[derive(Debug, Default, Clone, Copy)]
enum ManualDecodingState {
#[default]
ReadHeader,
ReadBody {
flags: u8,
size: usize,
},
}
#[derive(Debug)]
pub struct ZmtpManualParser {
state: ManualDecodingState,
max_msg_size: i64,
}
impl ZmtpManualParser {
pub fn new(max_msg_size: i64) -> Self {
Self {
state: ManualDecodingState::default(),
max_msg_size,
}
}
pub fn decode_frame_from_slice(&self, src: &[u8]) -> Result<Option<(Msg, usize)>, ZmqError> {
if src.len() < 2 {
return Ok(None);
}
let flags = src[0];
let is_long = (flags & ZMTP_FLAG_LONG) != 0;
let header_len = if is_long { 9 } else { 2 };
if src.len() < header_len {
return Ok(None);
}
let raw_size = if is_long {
let mut len_bytes = [0u8; 8];
len_bytes.copy_from_slice(&src[1..9]);
u64::from_be_bytes(len_bytes)
} else {
src[1] as u64
};
if self.max_msg_size >= 0 && raw_size > self.max_msg_size as u64 {
return Err(ZmqError::ProtocolViolation(format!(
"frame size {} exceeds ZMQ_MAXMSGSIZE {}",
raw_size, self.max_msg_size
)));
}
let size = raw_size as usize;
let total = header_len + size;
if src.len() < total {
return Ok(None);
}
let mut msg = Msg::from_vec(src[header_len..total].to_vec());
let mut rz_flags = MsgFlags::empty();
if (flags & ZMTP_FLAG_MORE) != 0 {
rz_flags |= MsgFlags::MORE;
}
if (flags & ZMTP_FLAG_COMMAND) != 0 {
rz_flags |= MsgFlags::COMMAND;
}
msg.set_flags(rz_flags);
Ok(Some((msg, total)))
}
pub fn peek_frame_len(&self, src: &[u8]) -> Result<Option<usize>, ZmqError> {
if src.is_empty() {
return Ok(None);
}
let flags = src[0];
let is_long = (flags & ZMTP_FLAG_LONG) != 0;
let header_len = if is_long { 9 } else { 2 };
if src.len() < header_len {
return Ok(None);
}
let raw_size = if is_long {
let mut len_bytes = [0u8; 8];
len_bytes.copy_from_slice(&src[1..9]);
u64::from_be_bytes(len_bytes)
} else {
src[1] as u64
};
if self.max_msg_size >= 0 && raw_size > self.max_msg_size as u64 {
return Err(ZmqError::ProtocolViolation(format!(
"frame size {} exceeds ZMQ_MAXMSGSIZE {}",
raw_size, self.max_msg_size
)));
}
Ok(Some(header_len + raw_size as usize))
}
pub fn decode_frame_from_bytes(&self, src: &Bytes) -> Result<Option<(Msg, usize)>, ZmqError> {
if src.len() < 2 {
return Ok(None);
}
let flags = src[0];
let is_long = (flags & ZMTP_FLAG_LONG) != 0;
let header_len = if is_long { 9 } else { 2 };
if src.len() < header_len {
return Ok(None);
}
let raw_size = if is_long {
let mut len_bytes = [0u8; 8];
len_bytes.copy_from_slice(&src[1..9]);
u64::from_be_bytes(len_bytes)
} else {
src[1] as u64
};
if self.max_msg_size >= 0 && raw_size > self.max_msg_size as u64 {
return Err(ZmqError::ProtocolViolation(format!(
"frame size {} exceeds ZMQ_MAXMSGSIZE {}",
raw_size, self.max_msg_size
)));
}
let size = raw_size as usize;
let total = header_len + size;
if src.len() < total {
return Ok(None);
}
let payload_bytes = src.slice(header_len..total);
let mut msg = Msg::from_bytes(payload_bytes);
let mut rz_flags = MsgFlags::empty();
if (flags & ZMTP_FLAG_MORE) != 0 {
rz_flags |= MsgFlags::MORE;
}
if (flags & ZMTP_FLAG_COMMAND) != 0 {
rz_flags |= MsgFlags::COMMAND;
}
msg.set_flags(rz_flags);
Ok(Some((msg, total)))
}
pub fn decode_from_buffer(&mut self, src: &mut BytesMut) -> Result<Option<Msg>, ZmqError> {
loop {
match self.state {
ManualDecodingState::ReadHeader => {
if src.is_empty() {
return Ok(None); }
let frame_flags_byte = src[0]; let is_long = (frame_flags_byte & ZMTP_FLAG_LONG) != 0;
let header_len = if is_long { 1 + 8 } else { 1 + 1 };
if src.len() < header_len {
return Ok(None); }
let raw_size = if is_long {
let mut len_bytes = [0u8; 8];
len_bytes.copy_from_slice(&src[1..9]);
u64::from_be_bytes(len_bytes)
} else {
src[1] as u64
};
if self.max_msg_size >= 0 && raw_size > self.max_msg_size as u64 {
return Err(ZmqError::ProtocolViolation(format!(
"frame size {} exceeds ZMQ_MAXMSGSIZE {}",
raw_size, self.max_msg_size
)));
}
let size = raw_size as usize;
if src.len() - header_len < size {
return Ok(None); }
let flags = src.get_u8();
let _ = src.split_to(header_len - 1); let body_bytes = src.split_to(size).freeze();
let mut msg = Msg::from_bytes(body_bytes);
let mut rz_flags = MsgFlags::empty();
if (flags & ZMTP_FLAG_MORE) != 0 {
rz_flags |= MsgFlags::MORE;
}
if (flags & ZMTP_FLAG_COMMAND) != 0 {
rz_flags |= MsgFlags::COMMAND;
}
msg.set_flags(rz_flags);
return Ok(Some(msg));
}
ManualDecodingState::ReadBody { flags, size } => {
if src.len() < size {
return Ok(None); }
let body_bytes = src.split_to(size).freeze();
self.state = ManualDecodingState::ReadHeader;
let mut msg = Msg::from_bytes(body_bytes);
let mut rz_flags = MsgFlags::empty();
if (flags & ZMTP_FLAG_MORE) != 0 {
rz_flags |= MsgFlags::MORE;
}
if (flags & ZMTP_FLAG_COMMAND) != 0 {
rz_flags |= MsgFlags::COMMAND;
}
msg.set_flags(rz_flags);
return Ok(Some(msg));
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn long_frame_header(size: u64) -> BytesMut {
let mut buf = BytesMut::new();
buf.put_u8(ZMTP_FLAG_LONG);
buf.put_u64(size);
buf
}
#[test]
fn oversized_frame_returns_protocol_violation() {
let mut parser = ZmtpManualParser::new(64);
let result = parser.decode_from_buffer(&mut long_frame_header(128));
assert!(matches!(result, Err(ZmqError::ProtocolViolation(_))));
}
#[test]
fn exact_limit_does_not_error() {
let mut parser = ZmtpManualParser::new(64);
let result = parser.decode_from_buffer(&mut long_frame_header(64));
assert!(matches!(result, Ok(None)));
}
#[test]
fn unlimited_allows_max_u64_header() {
let mut parser = ZmtpManualParser::new(-1);
let result = parser.decode_from_buffer(&mut long_frame_header(u64::MAX));
assert!(matches!(result, Ok(None)));
}
#[test]
fn zero_limit_blocks_any_nonzero_frame() {
let mut parser = ZmtpManualParser::new(0);
let result = parser.decode_from_buffer(&mut long_frame_header(1));
assert!(matches!(result, Err(ZmqError::ProtocolViolation(_))));
}
#[test]
fn peek_frame_len_short_frame() {
let parser = ZmtpManualParser::new(-1);
assert_eq!(parser.peek_frame_len(&[0x00, 5]).unwrap(), Some(7));
assert_eq!(
parser.peek_frame_len(&[0x00, 5, b'a', b'b']).unwrap(),
Some(7)
);
}
#[test]
fn peek_frame_len_long_frame() {
let parser = ZmtpManualParser::new(-1);
let mut hdr = vec![ZMTP_FLAG_LONG];
hdr.extend_from_slice(&300u64.to_be_bytes());
assert_eq!(parser.peek_frame_len(&hdr).unwrap(), Some(9 + 300));
}
#[test]
fn peek_frame_len_incomplete_header() {
let parser = ZmtpManualParser::new(-1);
assert_eq!(parser.peek_frame_len(&[]).unwrap(), None);
assert_eq!(parser.peek_frame_len(&[0x00]).unwrap(), None);
let mut partial = vec![ZMTP_FLAG_LONG];
partial.extend_from_slice(&[0, 0, 0]);
assert_eq!(parser.peek_frame_len(&partial).unwrap(), None);
}
#[test]
fn peek_frame_len_enforces_max_msg_size() {
let parser = ZmtpManualParser::new(64);
let mut hdr = vec![ZMTP_FLAG_LONG];
hdr.extend_from_slice(&128u64.to_be_bytes());
assert!(matches!(
parser.peek_frame_len(&hdr),
Err(ZmqError::ProtocolViolation(_))
));
}
#[test]
fn full_message_within_limit_decoded_correctly() {
let limit = 16i64;
let mut parser = ZmtpManualParser::new(limit);
let payload = b"hello rzmq";
let mut buf = BytesMut::new();
buf.put_u8(0x00); buf.put_u8(payload.len() as u8);
buf.put_slice(payload);
let result = parser.decode_from_buffer(&mut buf);
let msg = result
.expect("should succeed")
.expect("should have message");
assert_eq!(msg.data().unwrap(), payload);
}
}
#[cfg(test)]
mod additional_robustness_tests {
use super::*;
#[test]
fn test_fragmented_stream_parsing() {
let mut parser = ZmtpManualParser::new(1024);
let payload = b"stream-fragmentation-test-payload";
let mut raw_bytes = BytesMut::new();
raw_bytes.put_u8(0x00); raw_bytes.put_u8(payload.len() as u8);
raw_bytes.put_slice(payload);
let mut accumulator = BytesMut::new();
let total_len = raw_bytes.len();
for i in 0..total_len {
accumulator.put_u8(raw_bytes[i]);
let res = parser.decode_from_buffer(&mut accumulator);
if i < total_len - 1 {
assert!(
matches!(res, Ok(None)),
"Expected Ok(None) for incomplete frame at byte {}",
i
);
} else {
let msg = res
.expect("Should parse successfully on final byte")
.expect("Should yield a Msg on final byte");
assert_eq!(msg.data().unwrap(), payload.as_ref());
assert!(!msg.is_more());
}
}
}
#[test]
fn test_adversarial_header_parsing() {
let mut parser = ZmtpManualParser::new(1024);
let mut bad_buf = BytesMut::new();
bad_buf.put_u8(ZMTP_FLAG_COMMAND);
bad_buf.put_u8(0);
bad_buf.put_slice(b"garbage");
let res = parser.decode_from_buffer(&mut bad_buf);
assert!(res.is_ok());
}
#[test]
fn test_integer_overflow_protection() {
let mut parser = ZmtpManualParser::new(1024);
let mut bad_buf = BytesMut::new();
bad_buf.put_u8(ZMTP_FLAG_LONG);
bad_buf.put_u64(u64::MAX);
let res = parser.decode_from_buffer(&mut bad_buf);
assert!(
matches!(res, Err(ZmqError::ProtocolViolation(_))),
"Expected ProtocolViolation for u64::MAX frame size, got {:?}",
res
);
}
}