use std::io::Read;
use bytes::{Buf, BufMut, Bytes, BytesMut};
use crate::{TlvDecode, TlvEncode, TlvError, VarNum};
pub trait Tlv {
const TYP: usize;
fn inner_size(&self) -> usize;
fn critical() -> bool {
tlv_critical::<Self>()
}
fn from_reader(mut reader: impl Read) -> Result<Self, TlvError>
where
Self: TlvDecode,
{
let mut header_buf = [0; 18];
let bytes_read = reader.read(&mut header_buf).map_err(TlvError::IOError)?;
let mut header_bytes = Bytes::copy_from_slice(&header_buf);
let typ = VarNum::decode(&mut header_bytes)?;
if typ.value() as usize != Self::TYP {
return Err(TlvError::TypeMismatch {
expected: Self::TYP,
found: typ.value() as usize,
});
}
let len = VarNum::decode(&mut header_bytes)?;
let total_len = typ.size() + len.size() + len.value() as usize;
let mut bytes = BytesMut::with_capacity(total_len);
bytes.put(&header_buf[0..bytes_read]);
let mut left_to_read = total_len - bytes_read;
let mut buf = [0; 1024];
while left_to_read > 0 {
let bytes_read = reader
.read(&mut buf[0..left_to_read])
.map_err(TlvError::IOError)?;
bytes.put(&buf[..left_to_read]);
left_to_read -= bytes_read;
}
Self::decode(&mut bytes.freeze())
}
}
pub const fn tlv_critical<T: Tlv + ?Sized>() -> bool {
tlv_typ_critical(T::TYP)
}
pub const fn tlv_typ_critical(typ: usize) -> bool {
typ < 32 || typ & 1 == 1
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct GenericTlv<T> {
pub typ: VarNum,
pub len: VarNum,
pub content: T,
}
impl<T> TlvDecode for GenericTlv<T>
where
T: TlvDecode,
{
fn decode(bytes: &mut Bytes) -> crate::Result<Self> {
let typ = VarNum::decode(bytes)?.into();
let len = VarNum::decode(bytes)?;
if bytes.remaining() < len.into() {
return Err(TlvError::UnexpectedEndOfStream);
}
let mut inner_data = bytes.split_to(len.into());
Ok(Self {
typ,
len,
content: T::decode(&mut inner_data)?,
})
}
}
impl<T> TlvEncode for GenericTlv<T>
where
T: TlvEncode,
{
fn encode(&self) -> Bytes {
let mut bytes = BytesMut::with_capacity(self.size());
bytes.put(self.typ.encode());
bytes.put(self.len.encode());
let mut content = self.content.encode();
content.truncate(self.len.into());
if content.len() != self.len.into() {
panic!("GenericTLV length longer than encoded content");
}
bytes.put(content);
bytes.freeze()
}
fn size(&self) -> usize {
self.typ.size() + self.len.size() + self.content.size()
}
}
#[cfg(test)]
mod tests {
use bytes::{Buf, BufMut, Bytes, BytesMut};
use crate::tests::GenericNameComponent;
use crate::{error::TlvError, Result, TlvDecode, TlvEncode, VarNum};
use super::*;
#[derive(Debug)]
struct Name {
components: Vec<GenericNameComponent>,
}
impl Tlv for Name {
const TYP: usize = 7;
fn inner_size(&self) -> usize {
self.components.size()
}
}
impl TlvDecode for Name {
fn decode(mut bytes: &mut Bytes) -> Result<Self> {
let typ = VarNum::decode(&mut bytes)?;
if usize::from(typ) != Self::TYP {
return Err(TlvError::TypeMismatch {
expected: Self::TYP,
found: typ.into(),
});
}
let length = VarNum::decode(&mut bytes)?;
let mut inner_data = bytes.copy_to_bytes(length.into());
let components = Vec::<GenericNameComponent>::decode(&mut inner_data)?;
Ok(Self { components })
}
}
impl TlvEncode for Name {
fn encode(&self) -> Bytes {
let mut bytes = BytesMut::with_capacity(self.size());
bytes.put(VarNum::from(Self::TYP).encode());
bytes.put(VarNum::from(self.inner_size()).encode());
bytes.put(self.components.encode());
bytes.freeze()
}
fn size(&self) -> usize {
VarNum::from(Self::TYP).size()
+ VarNum::from(self.inner_size()).size()
+ self.components.size()
}
}
#[test]
fn wrong_type() {
let mut data = Bytes::from(&[9, 5, b'h', b'e', b'l', b'l', b'o', 255, 255, 255][..]);
let component = GenericNameComponent::decode(&mut data);
assert!(component.is_err());
let error = component.unwrap_err();
assert_eq!(
error,
TlvError::TypeMismatch {
expected: 8,
found: 9
}
);
}
#[test]
fn name() {
let mut data = Bytes::from(
&[
7, 14, 8, 5, b'h', b'e', b'l', b'l', b'o', 8, 5, b'w', b'o', b'r', b'l', b'd', 255,
255, 255,
][..],
);
let name = Name::decode(&mut data).unwrap();
assert_eq!(data.remaining(), 3);
assert_eq!(name.components.len(), 2);
assert_eq!(name.components[0].name, &b"hello"[..]);
assert_eq!(name.components[1].name, &b"world"[..]);
}
}