use crate::varint;
use std::fmt;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Error {
InvalidTag,
InvalidVarInt,
InvalidByteCount,
InvalidFinalSizeByte,
MissingBytes,
}
pub type Tag = u128;
pub mod tags {
use super::Tag;
pub const REPEAT: Tag = 0;
}
#[derive(Debug, Clone, PartialEq)]
pub struct Message {
pub tag: Tag,
pub body: Vec<u8>,
_private: bool,
}
impl Message {
pub fn new(tag: u128, body: Vec<u8>) -> Result<Self, Error> {
if tag == tags::REPEAT || tag > ((1 << 127) - 1) {
return Err(Error::InvalidTag);
}
if body.len() > u32::MAX as usize {
return Err(Error::InvalidByteCount);
}
Ok(Self {
tag,
body,
_private: false,
})
}
pub fn encode(messsages: Vec<Self>) -> Vec<u8> {
let mut bytes = Vec::new();
let len = messsages.len();
let mut last_tag = 0;
for (i, message) in messsages.into_iter().enumerate() {
let is_last = i == len - 1;
if message.tag == last_tag {
bytes.push(is_last as u8);
} else {
bytes.extend(varint::encode(2 * message.tag + (is_last as u128)));
last_tag = message.tag;
}
if !is_last {
bytes.extend(varint::encode(message.body.len() as u128));
}
bytes.extend(message.body);
}
bytes
}
pub fn decode(bytes: &[u8]) -> Result<Vec<Self>, Error> {
let mut messages = Vec::new();
let mut index = 0;
let mut last_tag = 0;
while index < bytes.len() {
let (value, size) =
varint::decode(&bytes[index..]).map_err(|_| Error::InvalidVarInt)?;
index += size;
let is_last = value % 2 == 1;
let tag = value / 2;
if tag == last_tag {
return Err(Error::InvalidTag);
}
let tag = if tag == tags::REPEAT {
last_tag
} else {
last_tag = tag;
tag
};
if is_last {
messages.push(Self {
tag,
body: bytes[index..].to_vec(),
_private: false,
});
break;
}
if index >= bytes.len() {
return Err(Error::MissingBytes);
}
let (n, size) = varint::decode(&bytes[index..]).map_err(|_| Error::InvalidVarInt)?;
index += size;
let length: usize = if n == 0 {
bytes.len() - index
} else if n > u32::MAX.into() {
return Err(Error::InvalidByteCount);
} else {
n.try_into().unwrap()
};
if index + length > bytes.len() {
return Err(Error::MissingBytes);
}
if n > 0 && index + length == bytes.len() {
return Err(Error::InvalidFinalSizeByte);
}
messages.push(Self {
tag,
body: bytes[index..(index + length)].to_vec(),
_private: false,
});
index += length;
}
Ok(messages)
}
}
impl std::error::Error for Error {}
impl fmt::Display for Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Error::InvalidTag => write!(f, "Invalid tag"),
Error::InvalidVarInt => write!(f, "Invalid variable integer encoding"),
Error::InvalidByteCount => write!(f, "Byte count exceeds 2^32 - 1"),
Error::InvalidFinalSizeByte => write!(f, "Final size byte must be zero"),
Error::MissingBytes => {
write!(f, "Variable-length encoding indicates bytes are missing")
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_new_valid() {
let data = Message::new(123, vec![1, 2, 3]).unwrap();
assert_eq!(data.tag, 123);
assert_eq!(data.body, vec![1, 2, 3]);
}
#[test]
fn test_new_invalid_tag() {
let result = Message::new(tags::REPEAT, vec![1, 2, 3]);
assert_eq!(result.err(), Some(Error::InvalidTag));
let result = Message::new(1 << 127, vec![1, 2, 3]);
assert_eq!(result.err(), Some(Error::InvalidTag));
}
#[test]
fn test_new_invalid_byte_count() {
if u32::MAX as usize == usize::MAX {
return;
}
let result = Message::new(1, vec![0; (u32::MAX as usize) + 1]);
assert_eq!(result.err(), Some(Error::InvalidByteCount));
}
#[test]
fn test_encode_single_chunk() {
let chunk = Message::new(1, vec![5, 6, 7]).unwrap();
let encoded = Message::encode(vec![chunk]);
assert_eq!(encoded, vec![3, 5, 6, 7]);
}
#[test]
fn test_encode_multiple_chunks() {
let chunk1 = Message::new(1, vec![1, 2]).unwrap();
let chunk2 = Message::new(2, vec![3, 4, 5]).unwrap();
let encoded = Message::encode(vec![chunk1, chunk2]);
assert_eq!(encoded, vec![2, 2, 1, 2, 5, 3, 4, 5]);
}
#[test]
fn test_encode_repeated_tag() {
let chunk1 = Message::new(1, vec![1, 2]).unwrap();
let chunk2 = Message::new(1, vec![3, 4]).unwrap();
let chunk3 = Message::new(2, vec![5, 6]).unwrap();
let encoded = Message::encode(vec![chunk1, chunk2, chunk3]);
assert_eq!(encoded, vec![2, 2, 1, 2, 0, 2, 3, 4, 5, 5, 6]);
}
#[test]
fn test_decode_valid() {
let encoded = vec![2, 2, 1, 2, 5, 3, 4, 5];
let decoded = Message::decode(&encoded).unwrap();
assert_eq!(decoded.len(), 2);
assert_eq!(decoded[0].tag, 1);
assert_eq!(decoded[0].body, vec![1, 2]);
assert_eq!(decoded[1].tag, 2);
assert_eq!(decoded[1].body, vec![3, 4, 5]);
}
#[test]
fn test_decode_repeated_tag() {
let encoded = vec![2, 2, 1, 2, 0, 2, 3, 4, 5, 5, 6];
let decoded = Message::decode(&encoded).unwrap();
assert_eq!(decoded.len(), 3);
assert_eq!(decoded[0].tag, 1);
assert_eq!(decoded[0].body, vec![1, 2]);
assert_eq!(decoded[1].tag, 1); assert_eq!(decoded[1].body, vec![3, 4]);
assert_eq!(decoded[2].tag, 2);
assert_eq!(decoded[2].body, vec![5, 6]);
}
#[test]
fn test_decode_invalid_varint() {
let encoded = vec![0xFF];
let result = Message::decode(&encoded);
assert_eq!(result.err(), Some(Error::InvalidVarInt));
}
#[test]
fn test_decode_missing_bytes() {
let encoded = vec![2, 10, 1, 2, 3];
let result = Message::decode(&encoded);
assert_eq!(result.err(), Some(Error::MissingBytes));
}
#[test]
fn test_decode_invalid_final_size_byte() {
let encoded = vec![2, 3, 1, 2, 3];
let result = Message::decode(&encoded);
assert_eq!(result.err(), Some(Error::InvalidFinalSizeByte));
}
#[test]
fn test_empty_final_body() {
let chunk = Message::new(3, vec![]).unwrap();
let encoded = Message::encode(vec![chunk]);
assert_eq!(encoded, vec![7]);
let decoded = Message::decode(&encoded).unwrap();
assert_eq!(decoded.len(), 1);
assert_eq!(decoded[0].tag, 3);
assert_eq!(decoded[0].body, Vec::<u8>::new());
}
#[test]
fn test_decode_with_termination_only() {
let encoded = vec![11];
let decoded = Message::decode(&encoded).unwrap();
assert_eq!(decoded.len(), 1);
assert_eq!(decoded[0].tag, 5);
assert_eq!(decoded[0].body, Vec::<u8>::new());
}
#[test]
fn test_roundtrip_encode_decode() {
let original = vec![
Message::new(1, vec![1, 2]).unwrap(),
Message::new(2, vec![3, 4, 5]).unwrap(),
];
let encoded = Message::encode(original.clone());
let decoded = Message::decode(&encoded).unwrap();
assert_eq!(original, decoded);
}
#[test]
fn test_large_tag_values() {
let large_tag = (1 << 127) - 1;
let data = Message::new(large_tag, vec![9, 8, 7]).unwrap();
let encoded = Message::encode(vec![data]);
let decoded = Message::decode(&encoded).unwrap();
assert_eq!(decoded.len(), 1);
assert_eq!(decoded[0].tag, large_tag);
assert_eq!(decoded[0].body, vec![9, 8, 7]);
}
#[test]
fn test_multi_chunk_with_repeated_tag() {
let chunks = vec![
Message::new(10, vec![1, 2]).unwrap(),
Message::new(10, vec![3, 4]).unwrap(),
Message::new(10, vec![5, 6]).unwrap(),
Message::new(20, vec![7, 8, 9]).unwrap(),
];
let encoded = Message::encode(chunks.clone());
let decoded = Message::decode(&encoded).unwrap();
assert_eq!(chunks, decoded);
}
}