use super::{Decode, DecodeError, Encode, EncodeError, take};
pub fn encode_counted_u8<T: Encode>(items: &[T], out: &mut Vec<u8>) -> Result<(), EncodeError> {
let count: u8 = items
.len()
.try_into()
.map_err(|_| EncodeError::CountOverflow {
length: items.len(),
limit: u8::MAX.into(),
})?;
out.push(count);
encode_items(items, out)
}
pub fn encode_counted_u32<T: Encode>(items: &[T], out: &mut Vec<u8>) -> Result<(), EncodeError> {
let count: u32 = items
.len()
.try_into()
.map_err(|_| EncodeError::CountOverflow {
length: items.len(),
limit: u32::MAX.into(),
})?;
count.encode(out)?;
encode_items(items, out)
}
pub fn decode_counted_u8<T: Decode>(input: &mut &[u8]) -> Result<Vec<T>, DecodeError> {
let count = u8::decode(input)? as usize;
decode_items(input, count)
}
pub fn decode_counted_u32<T: Decode>(input: &mut &[u8]) -> Result<Vec<T>, DecodeError> {
let count = u32::decode(input)? as usize;
decode_items(input, count)
}
pub fn encode_bytes_u8(bytes: &[u8], out: &mut Vec<u8>) -> Result<(), EncodeError> {
let length: u8 = bytes
.len()
.try_into()
.map_err(|_| EncodeError::BytesOverflow {
length: bytes.len(),
limit: u8::MAX.into(),
})?;
out.push(length);
out.extend_from_slice(bytes);
Ok(())
}
pub fn encode_bytes_u32(bytes: &[u8], out: &mut Vec<u8>) -> Result<(), EncodeError> {
let length: u32 = bytes
.len()
.try_into()
.map_err(|_| EncodeError::BytesOverflow {
length: bytes.len(),
limit: u32::MAX.into(),
})?;
length.encode(out)?;
out.extend_from_slice(bytes);
Ok(())
}
pub fn decode_bytes_u8(input: &mut &[u8]) -> Result<Vec<u8>, DecodeError> {
let length = u8::decode(input)? as usize;
Ok(take(input, length)?.to_vec())
}
pub fn decode_bytes_u32(input: &mut &[u8]) -> Result<Vec<u8>, DecodeError> {
let length = u32::decode(input)? as usize;
Ok(take(input, length)?.to_vec())
}
fn encode_items<T: Encode>(items: &[T], out: &mut Vec<u8>) -> Result<(), EncodeError> {
for item in items {
item.encode(out)?;
}
Ok(())
}
const MAX_PREALLOCATION_BYTES: usize = 64 * 1024;
fn decode_items<T: Decode>(input: &mut &[u8], count: usize) -> Result<Vec<T>, DecodeError> {
let capacity = count
.min(input.len())
.min(MAX_PREALLOCATION_BYTES / size_of::<T>().max(1));
let mut items = Vec::with_capacity(capacity);
for _ in 0..count {
items.push(T::decode(input)?);
}
Ok(items)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn counted_u8_round_trips() {
let items: Vec<u32> = vec![1, 2];
let mut out = Vec::new();
encode_counted_u8(&items, &mut out).unwrap();
assert_eq!(out, [2, 1, 0, 0, 0, 2, 0, 0, 0]);
let mut input = out.as_slice();
assert_eq!(decode_counted_u8::<u32>(&mut input).unwrap(), items);
assert!(input.is_empty());
}
#[test]
fn counted_u32_round_trips() {
let items: Vec<u8> = vec![9, 8, 7];
let mut out = Vec::new();
encode_counted_u32(&items, &mut out).unwrap();
assert_eq!(out, [3, 0, 0, 0, 9, 8, 7]);
let mut input = out.as_slice();
assert_eq!(decode_counted_u32::<u8>(&mut input).unwrap(), items);
assert!(input.is_empty());
}
#[test]
fn counted_u8_overflows_at_256_elements() {
let items = vec![0u8; 256];
let mut out = Vec::new();
let error = encode_counted_u8(&items, &mut out).unwrap_err();
assert!(matches!(
error,
EncodeError::CountOverflow {
length: 256,
limit: 255,
}
));
}
#[test]
fn truncated_counted_collection_fails_to_decode() {
let mut input: &[u8] = &[3, 1, 2];
let error = decode_counted_u8::<u8>(&mut input).unwrap_err();
assert!(matches!(error, DecodeError::UnexpectedEof));
}
#[test]
fn hostile_count_does_not_allocate() {
let mut input: &[u8] = &[0xff, 0xff, 0xff, 0xff, 1];
let error = decode_counted_u32::<u8>(&mut input).unwrap_err();
assert!(matches!(error, DecodeError::UnexpectedEof));
}
#[test]
fn bytes_u8_round_trips() {
let mut out = Vec::new();
encode_bytes_u8(&[1, 2, 3], &mut out).unwrap();
assert_eq!(out, [3, 1, 2, 3]);
let mut input = out.as_slice();
assert_eq!(decode_bytes_u8(&mut input).unwrap(), vec![1, 2, 3]);
}
#[test]
fn bytes_u8_overflows_at_256_bytes() {
let mut out = Vec::new();
let error = encode_bytes_u8(&[0; 256], &mut out).unwrap_err();
assert!(matches!(error, EncodeError::BytesOverflow { .. }));
}
#[test]
fn bytes_u32_round_trips() {
let mut out = Vec::new();
encode_bytes_u32(&[5, 6], &mut out).unwrap();
assert_eq!(out, [2, 0, 0, 0, 5, 6]);
let mut input = out.as_slice();
assert_eq!(decode_bytes_u32(&mut input).unwrap(), vec![5, 6]);
}
}