use thiserror::Error;
pub const SPECIAL: u8 = 0;
pub const RELAY: u8 = 1;
pub const AUTHOR: u8 = 2;
pub const KIND: u8 = 3;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
#[non_exhaustive]
pub enum TlvError {
#[error("TLV value is too long for a 1-byte length field: {len} bytes (max 255)")]
ValueTooLong {
len: usize,
},
#[error("TLV record truncated at offset {offset}: missing length byte")]
TruncatedHeader {
offset: usize,
},
#[error(
"TLV record truncated at offset {offset}: expected {expected} value bytes, only {available} remain"
)]
TruncatedValue {
offset: usize,
expected: usize,
available: usize,
},
}
#[derive(Debug, Clone, Copy)]
pub struct Record<'a> {
pub tag: u8,
pub value: &'a [u8],
}
pub fn encode<'a, I>(records: I) -> Result<Vec<u8>, TlvError>
where
I: IntoIterator<Item = (u8, &'a [u8])>,
{
let records = records.into_iter();
let (lower, _) = records.size_hint();
let mut out = Vec::with_capacity(lower * 2);
for (tag, value) in records {
let len =
u8::try_from(value.len()).map_err(|_| TlvError::ValueTooLong { len: value.len() })?;
out.push(tag);
out.push(len);
out.extend_from_slice(value);
}
Ok(out)
}
#[must_use]
pub const fn iter(bytes: &[u8]) -> RecordIter<'_> {
RecordIter {
bytes,
cursor: 0,
finished: false,
}
}
#[derive(Debug)]
pub struct RecordIter<'a> {
bytes: &'a [u8],
cursor: usize,
finished: bool,
}
impl<'a> Iterator for RecordIter<'a> {
type Item = Result<Record<'a>, TlvError>;
fn next(&mut self) -> Option<Self::Item> {
if self.finished {
return None;
}
let offset = self.cursor;
let remaining = self.bytes.get(offset..)?;
let (&tag, after_tag) = remaining.split_first()?;
let Some((&len_byte, value_buf)) = after_tag.split_first() else {
self.finished = true;
return Some(Err(TlvError::TruncatedHeader { offset }));
};
let len = len_byte as usize;
if value_buf.len() < len {
self.finished = true;
return Some(Err(TlvError::TruncatedValue {
offset: offset + 2,
expected: len,
available: value_buf.len(),
}));
}
let (value, _) = value_buf.split_at(len);
self.cursor = offset + 2 + len;
Some(Ok(Record { tag, value }))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trip_single_record() {
let payload = [0xab; 32];
let encoded = encode([(SPECIAL, payload.as_slice())]).unwrap();
let records: Result<Vec<_>, _> = iter(&encoded).collect();
let records = records.unwrap();
assert_eq!(records.len(), 1);
assert_eq!(records[0].tag, SPECIAL);
assert_eq!(records[0].value, payload);
}
#[test]
fn round_trip_multiple_records() {
let pubkey = [0x01; 32];
let relay: &[u8] = b"wss://relay.example";
let kind = [0x00, 0x00, 0x00, 0x01];
let encoded = encode([
(SPECIAL, pubkey.as_slice()),
(RELAY, relay),
(KIND, kind.as_slice()),
])
.unwrap();
let records: Vec<_> = iter(&encoded).map(Result::unwrap).collect();
assert_eq!(records.len(), 3);
assert_eq!(records[0].value, pubkey);
assert_eq!(records[1].value, relay);
assert_eq!(records[2].value, kind);
}
#[test]
fn encode_rejects_oversized_value() {
let big = vec![0u8; 300];
let err = encode([(SPECIAL, big.as_slice())]).unwrap_err();
assert!(matches!(err, TlvError::ValueTooLong { len: 300 }));
}
#[test]
fn truncated_header_is_rejected() {
let bytes = [SPECIAL]; let err = iter(&bytes).next().unwrap().unwrap_err();
assert!(matches!(err, TlvError::TruncatedHeader { offset: 0 }));
}
#[test]
fn truncated_value_is_rejected() {
let bytes = [SPECIAL, 0x05, 0x00, 0x00];
let err = iter(&bytes).next().unwrap().unwrap_err();
assert!(matches!(
err,
TlvError::TruncatedValue {
expected: 5,
available: 2,
..
}
));
}
#[test]
fn empty_input_produces_no_records() {
assert!(iter(&[]).next().is_none());
}
}