use heapless::Vec;
use crate::limits::MAX_TLVS;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Tlv<'a> {
pub tag: u32,
pub value: &'a [u8],
}
const TAG_MULTIBYTE: u8 = 0x1F;
const TAG_MORE: u8 = 0x80;
const TAG_MAX_OCTETS: usize = 4;
const LEN_LONG: u8 = 0x80;
pub fn parse(input: &[u8]) -> Result<Vec<Tlv<'_>, MAX_TLVS>, TlvError> {
let mut out = Vec::new();
let mut pos = 0usize;
while pos < input.len() {
let tag = parse_tag(input, &mut pos)?;
let len = parse_len(input, &mut pos)?;
let end = pos.checked_add(len).ok_or(TlvError::Truncated)?;
if end > input.len() {
return Err(TlvError::Truncated);
}
out.push(Tlv {
tag,
value: &input[pos..end],
})
.map_err(|_| TlvError::TooMany)?;
pos = end;
}
Ok(out)
}
pub fn encode(items: &[Tlv<'_>], out: &mut [u8]) -> Result<usize, TlvError> {
let mut pos = 0usize;
for item in items {
pos = write_tag(item.tag, out, pos)?;
pos = write_len(item.value.len(), out, pos)?;
pos = write_bytes(out, pos, item.value)?;
}
Ok(pos)
}
fn parse_tag(input: &[u8], pos: &mut usize) -> Result<u32, TlvError> {
let b0 = *input.get(*pos).ok_or(TlvError::Truncated)?;
*pos += 1;
let mut tag = u32::from(b0);
if b0 & TAG_MULTIBYTE == TAG_MULTIBYTE {
let mut octets = 1usize;
loop {
let b = *input.get(*pos).ok_or(TlvError::Truncated)?;
*pos += 1;
octets += 1;
if octets > TAG_MAX_OCTETS {
return Err(TlvError::BadLength);
}
tag = (tag << 8) | u32::from(b);
if b & TAG_MORE == 0 {
break;
}
}
}
Ok(tag)
}
fn parse_len(input: &[u8], pos: &mut usize) -> Result<usize, TlvError> {
let b0 = *input.get(*pos).ok_or(TlvError::Truncated)?;
*pos += 1;
if b0 & LEN_LONG == 0 {
return Ok(usize::from(b0));
}
let count = usize::from(b0 & !LEN_LONG);
if count == 0 || count > core::mem::size_of::<usize>() {
return Err(TlvError::BadLength);
}
let mut len = 0usize;
for _ in 0..count {
let b = *input.get(*pos).ok_or(TlvError::Truncated)?;
*pos += 1;
len = (len << 8) | usize::from(b);
}
Ok(len)
}
fn write_tag(tag: u32, out: &mut [u8], pos: usize) -> Result<usize, TlvError> {
let bytes = tag.to_be_bytes();
let start = bytes
.iter()
.position(|&b| b != 0)
.unwrap_or(bytes.len() - 1);
write_bytes(out, pos, &bytes[start..])
}
fn write_len(len: usize, out: &mut [u8], pos: usize) -> Result<usize, TlvError> {
if len < usize::from(LEN_LONG) {
let b = u8::try_from(len).map_err(|_| TlvError::Overflow)?;
return write_bytes(out, pos, &[b]);
}
let bytes = len.to_be_bytes();
let start = bytes
.iter()
.position(|&b| b != 0)
.unwrap_or(bytes.len() - 1);
let body = &bytes[start..];
let count = u8::try_from(body.len()).map_err(|_| TlvError::Overflow)?;
let pos = write_bytes(out, pos, &[LEN_LONG | count])?;
write_bytes(out, pos, body)
}
fn write_bytes(out: &mut [u8], pos: usize, src: &[u8]) -> Result<usize, TlvError> {
let end = pos.checked_add(src.len()).ok_or(TlvError::Overflow)?;
if end > out.len() {
return Err(TlvError::Overflow);
}
out[pos..end].copy_from_slice(src);
Ok(end)
}
#[derive(thiserror::Error, Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum TlvError {
#[error("TLV value runs past end of buffer")]
Truncated,
#[error("malformed TLV length encoding")]
BadLength,
#[error("more than MAX_TLVS top-level TLV objects")]
TooMany,
#[error("encode output buffer too small")]
Overflow,
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
use scll_test_util::HexSlice;
#[test]
fn parses_two_single_byte_tag_objects() {
let input = [0x4F, 0x02, 0xA0, 0x00, 0x9F, 0x70, 0x01, 0x07];
let tlvs = parse(&input).unwrap();
assert_eq!(tlvs.len(), 2);
assert_eq!(tlvs[0].tag, 0x4F);
assert_eq!(HexSlice(tlvs[0].value), HexSlice(&[0xA0, 0x00]));
assert_eq!(tlvs[1].tag, 0x9F70);
assert_eq!(HexSlice(tlvs[1].value), HexSlice(&[0x07]));
}
#[test]
fn parses_long_form_length() {
let mut input = heapless::Vec::<u8, 300>::new();
input.extend_from_slice(&[0x66, 0x81, 0x82]).unwrap();
input.extend_from_slice(&[0xAB; 0x82]).unwrap();
let tlvs = parse(&input).unwrap();
assert_eq!(tlvs.len(), 1);
assert_eq!(tlvs[0].tag, 0x66);
assert_eq!(tlvs[0].value.len(), 0x82);
}
#[test]
fn empty_input_yields_no_objects() {
assert_eq!(parse(&[]).unwrap().len(), 0);
}
#[test]
fn truncated_value_is_rejected() {
assert_eq!(parse(&[0x4F, 0x05, 0x01, 0x02]), Err(TlvError::Truncated));
}
#[test]
fn truncated_tag_is_rejected() {
assert_eq!(parse(&[0x9F]), Err(TlvError::Truncated));
}
#[test]
fn missing_length_octet_is_rejected() {
assert_eq!(parse(&[0x4F]), Err(TlvError::Truncated));
}
#[test]
fn indefinite_length_is_rejected() {
assert_eq!(parse(&[0x4F, 0x80, 0x01]), Err(TlvError::BadLength));
}
#[test]
fn oversized_tag_is_rejected() {
assert_eq!(
parse(&[0x1F, 0x81, 0x81, 0x81, 0x81, 0x01]),
Err(TlvError::BadLength)
);
}
#[test]
fn long_form_length_octets_exceeding_usize_rejected() {
assert_eq!(parse(&[0x4F, 0x89]), Err(TlvError::BadLength));
}
#[test]
fn too_many_objects_is_rejected() {
let mut input = heapless::Vec::<u8, { (MAX_TLVS + 1) * 3 }>::new();
for _ in 0..=MAX_TLVS {
input.extend_from_slice(&[0x80, 0x01, 0x00]).unwrap();
}
assert_eq!(parse(&input), Err(TlvError::TooMany));
}
#[test]
fn encode_emits_canonical_long_form() {
let value = [0xAB; 0x82];
let items = [Tlv {
tag: 0x66,
value: &value,
}];
let mut out = [0u8; 300];
let n = encode(&items, &mut out).unwrap();
assert_eq!(HexSlice(&out[..3]), HexSlice(&[0x66, 0x81, 0x82]));
assert_eq!(n, 3 + 0x82);
}
#[test]
fn encode_overflow_is_reported_not_panicked() {
let items = [Tlv {
tag: 0x4F,
value: &[1, 2, 3, 4],
}];
let mut out = [0u8; 3]; assert_eq!(encode(&items, &mut out), Err(TlvError::Overflow));
}
fn tag_strategy() -> impl Strategy<Value = u32> {
prop_oneof![
(0u32..=0xFF).prop_filter("not a continuation tag", |t| t & 0x1F != 0x1F),
(0u32..=0x7F).prop_map(|lo| 0x9F00 | lo),
]
}
fn tlv_items() -> impl Strategy<Value = std::vec::Vec<(u32, std::vec::Vec<u8>)>> {
proptest::collection::vec(
(
tag_strategy(),
proptest::collection::vec(any::<u8>(), 0..=64),
),
0..=MAX_TLVS,
)
}
proptest! {
#[test]
fn parse_after_encode_is_identity(items in tlv_items()) {
let tlvs: std::vec::Vec<Tlv> =
items.iter().map(|(t, v)| Tlv { tag: *t, value: v }).collect();
let mut out = [0u8; 64 * 70];
let n = encode(&tlvs, &mut out).unwrap();
let parsed = parse(&out[..n]).unwrap();
prop_assert_eq!(parsed.len(), tlvs.len());
for (got, want) in parsed.iter().zip(tlvs.iter()) {
prop_assert_eq!(got.tag, want.tag);
prop_assert_eq!(got.value, want.value);
}
}
#[test]
fn encoder_output_always_parses(items in tlv_items()) {
let tlvs: std::vec::Vec<Tlv> =
items.iter().map(|(t, v)| Tlv { tag: *t, value: v }).collect();
let mut out = [0u8; 64 * 70];
let n = encode(&tlvs, &mut out).unwrap();
prop_assert!(parse(&out[..n]).is_ok());
}
#[test]
fn parse_arbitrary_never_panics(bytes in proptest::collection::vec(any::<u8>(), 0..512)) {
let _ = parse(&bytes);
}
}
}