use crate::Bytes;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DigMessage {
pub msg_type: u8,
pub id: Option<u16>,
pub data: Bytes,
}
impl DigMessage {
pub const MAX_MESSAGE_SIZE: usize = 16 * 1024 * 1024;
pub fn new(msg_type: u8, id: Option<u16>, data: Bytes) -> Self {
Self { msg_type, id, data }
}
pub fn to_bytes(&self) -> Vec<u8> {
let mut buf = Vec::with_capacity(1 + 1 + 2 + 4 + self.data.len());
buf.push(self.msg_type);
match self.id {
Some(id) => {
buf.push(1); buf.extend_from_slice(&id.to_be_bytes());
}
None => {
buf.push(0); }
}
buf.extend_from_slice(&(self.data.len() as u32).to_be_bytes());
buf.extend_from_slice(&self.data);
buf
}
pub fn from_bytes(bytes: &[u8]) -> Option<Self> {
if bytes.len() < 2 {
return None;
}
let msg_type = bytes[0];
let has_id = bytes[1];
let mut offset: usize = 2;
let id = if has_id != 0 {
let after_id = offset.checked_add(2)?;
if bytes.len() < after_id {
return None;
}
let id = u16::from_be_bytes([bytes[offset], bytes[offset + 1]]);
offset = after_id;
Some(id)
} else {
None
};
let after_len = offset.checked_add(4)?;
if bytes.len() < after_len {
return None;
}
let data_len = u32::from_be_bytes([
bytes[offset],
bytes[offset + 1],
bytes[offset + 2],
bytes[offset + 3],
]) as usize;
offset = after_len;
if data_len > Self::MAX_MESSAGE_SIZE {
return None;
}
let end = offset.checked_add(data_len)?;
if bytes.len() < end {
return None;
}
let data = Bytes::new(bytes[offset..end].to_vec());
Some(Self { msg_type, id, data })
}
pub fn from_bytes_owned(mut buf: Vec<u8>) -> Option<Self> {
if buf.len() < 2 {
return None;
}
let msg_type = buf[0];
let has_id = buf[1];
let mut offset: usize = 2;
let id = if has_id != 0 {
let after_id = offset.checked_add(2)?;
if buf.len() < after_id {
return None;
}
let id = u16::from_be_bytes([buf[offset], buf[offset + 1]]);
offset = after_id;
Some(id)
} else {
None
};
let after_len = offset.checked_add(4)?;
if buf.len() < after_len {
return None;
}
let data_len = u32::from_be_bytes([
buf[offset],
buf[offset + 1],
buf[offset + 2],
buf[offset + 3],
]) as usize;
offset = after_len;
if data_len > Self::MAX_MESSAGE_SIZE {
return None;
}
let end = offset.checked_add(data_len)?;
if buf.len() < end {
return None;
}
let data: Vec<u8> = buf.drain(offset..end).collect();
Some(Self {
msg_type,
id,
data: Bytes::new(data),
})
}
pub fn is_dig_extension(&self) -> bool {
self.msg_type >= 200
}
pub fn is_chia_standard(&self) -> bool {
self.msg_type < 200
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trip_chia_opcode() {
let msg = DigMessage::new(20, Some(42), Bytes::new(vec![1, 2, 3]));
let wire = msg.to_bytes();
let decoded = DigMessage::from_bytes(&wire).expect("decode");
assert_eq!(decoded.msg_type, 20);
assert_eq!(decoded.id, Some(42));
assert_eq!(decoded.data.as_ref(), &[1, 2, 3]);
}
#[test]
fn round_trip_dig_opcode() {
let msg = DigMessage::new(218, None, Bytes::new(vec![0xAB]));
let wire = msg.to_bytes();
let decoded = DigMessage::from_bytes(&wire).expect("decode");
assert_eq!(decoded.msg_type, 218);
assert_eq!(decoded.id, None);
assert!(decoded.is_dig_extension());
assert!(!decoded.is_chia_standard());
}
#[test]
fn from_bytes_owned_avoids_copy_when_buffer_is_owned() {
let msg = DigMessage::new(218, Some(7), Bytes::new(vec![9, 9, 9]));
let wire: Vec<u8> = msg.to_bytes();
let decoded = DigMessage::from_bytes_owned(wire).expect("decode");
assert_eq!(decoded, msg);
}
#[test]
fn from_bytes_owned_rejects_same_malformed_inputs_as_from_bytes() {
assert!(DigMessage::from_bytes_owned(vec![]).is_none());
assert!(DigMessage::from_bytes_owned(vec![20, 1]).is_none());
let mut wire = vec![1u8, 0];
wire.extend_from_slice(&((DigMessage::MAX_MESSAGE_SIZE as u32) + 1).to_be_bytes());
assert!(DigMessage::from_bytes_owned(wire).is_none());
let truncated = vec![20u8, 0, 0x00, 0x00, 0x00, 0x04, 0xAB, 0xCD];
assert!(DigMessage::from_bytes_owned(truncated).is_none());
}
#[test]
fn from_bytes_owned_ignores_trailing_bytes_like_from_bytes() {
let mut wire = vec![20u8, 0, 0x00, 0x00, 0x00, 0x02, 0xAB, 0xCD];
wire.extend_from_slice(&[0xEE, 0xFF]); let decoded = DigMessage::from_bytes_owned(wire).expect("decode");
assert_eq!(decoded.data.as_ref(), &[0xAB, 0xCD]);
}
#[test]
fn empty_buffer_returns_none() {
assert!(DigMessage::from_bytes(&[]).is_none());
assert!(DigMessage::from_bytes(&[20]).is_none());
}
#[test]
fn truncated_id_prefix_returns_none() {
assert!(DigMessage::from_bytes(&[20, 1]).is_none());
assert!(DigMessage::from_bytes(&[20, 1, 0x00]).is_none());
}
#[test]
fn truncated_data_len_prefix_returns_none() {
assert!(DigMessage::from_bytes(&[20, 0]).is_none());
assert!(DigMessage::from_bytes(&[20, 0, 0x00, 0x00, 0x00]).is_none());
assert!(DigMessage::from_bytes(&[20, 1, 0x00, 0x2A, 0x00, 0x00]).is_none());
}
#[test]
fn truncated_data_returns_none() {
let wire = [20u8, 0, 0x00, 0x00, 0x00, 0x04, 0xAB, 0xCD];
assert!(DigMessage::from_bytes(&wire).is_none());
let ok = [20u8, 0, 0x00, 0x00, 0x00, 0x02, 0xAB, 0xCD];
let decoded = DigMessage::from_bytes(&ok).expect("exact-length payload decodes");
assert_eq!(decoded.data.as_ref(), &[0xAB, 0xCD]);
}
#[test]
fn oversized_data_len_is_rejected() {
let over = (DigMessage::MAX_MESSAGE_SIZE as u32) + 1;
let mut wire = vec![20u8, 0]; wire.extend_from_slice(&over.to_be_bytes());
assert!(DigMessage::from_bytes(&wire).is_none());
}
#[test]
fn max_data_len_prefix_does_not_panic_and_is_rejected() {
let mut wire = vec![20u8, 0]; wire.extend_from_slice(&u32::MAX.to_be_bytes());
assert!(DigMessage::from_bytes(&wire).is_none());
}
#[test]
fn offset_plus_data_len_overflow_is_checked_not_wrapping() {
let bytes = [20u8, 0, 0, 0, 0, 0]; assert!(DigMessage::from_bytes(&bytes).is_some());
let mut wire = vec![1u8, 0];
wire.extend_from_slice(&u32::MAX.to_be_bytes());
let result = std::panic::catch_unwind(|| DigMessage::from_bytes(&wire));
assert!(
result.is_ok(),
"from_bytes must not panic on data_len = u32::MAX"
);
assert_eq!(result.unwrap(), None);
}
#[test]
fn data_len_one_byte_over_cap_is_rejected_at_exact_boundary() {
let over_by_one = (DigMessage::MAX_MESSAGE_SIZE as u32) + 1;
let mut wire = vec![1u8, 0];
wire.extend_from_slice(&over_by_one.to_be_bytes());
assert!(DigMessage::from_bytes(&wire).is_none());
}
#[test]
fn zero_length_data_round_trip() {
let msg = DigMessage::new(64, Some(1), Bytes::default());
let wire = msg.to_bytes();
let decoded = DigMessage::from_bytes(&wire).expect("zero-length decode");
assert_eq!(decoded, msg);
assert!(decoded.data.as_ref().is_empty());
}
#[test]
fn dig_extension_boundary_at_200() {
let below = DigMessage::new(199, None, Bytes::default());
assert!(below.is_chia_standard());
assert!(!below.is_dig_extension());
let at = DigMessage::new(200, None, Bytes::default());
assert!(at.is_dig_extension());
assert!(!at.is_chia_standard());
}
}