use alloc::vec::Vec;
use crate::error::{Error, Result};
pub const FIXED_HET_MIN: u8 = 128;
pub const WORD: usize = 4;
#[derive(Debug, Clone, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct HeaderExtension<'a> {
pub het: u8,
pub content: &'a [u8],
}
impl<'a> HeaderExtension<'a> {
pub fn new(het: u8, content: &'a [u8]) -> Self {
HeaderExtension { het, content }
}
pub fn is_fixed(&self) -> bool {
self.het >= FIXED_HET_MIN
}
pub fn serialized_len(&self) -> usize {
if self.is_fixed() {
WORD
} else {
2 + self.content.len()
}
}
pub fn hel(&self) -> usize {
self.serialized_len() / WORD
}
pub fn parse(data: &'a [u8]) -> Result<(Self, usize)> {
if data.is_empty() {
return Err(Error::BufferTooShort {
need: 1,
have: 0,
what: "header extension HET",
});
}
let het = data[0];
if het >= FIXED_HET_MIN {
if data.len() < WORD {
return Err(Error::BufferTooShort {
need: WORD,
have: data.len(),
what: "fixed-length header extension",
});
}
Ok((
HeaderExtension {
het,
content: &data[1..WORD],
},
WORD,
))
} else {
if data.len() < 2 {
return Err(Error::BufferTooShort {
need: 2,
have: data.len(),
what: "variable-length header extension HEL",
});
}
let hel = data[1] as usize;
if hel == 0 {
return Err(Error::InvalidExtension {
reason: "HEL must be >= 1 for a variable-length extension",
});
}
let total = hel * WORD;
if data.len() < total {
return Err(Error::BufferTooShort {
need: total,
have: data.len(),
what: "variable-length header extension content",
});
}
Ok((
HeaderExtension {
het,
content: &data[2..total],
},
total,
))
}
}
pub fn serialize_into(&self, out: &mut [u8]) -> Result<usize> {
let total = self.serialized_len();
if out.len() < total {
return Err(Error::OutputBufferTooSmall {
need: total,
have: out.len(),
});
}
out[0] = self.het;
if self.is_fixed() {
if self.content.len() != WORD - 1 {
return Err(Error::InvalidExtension {
reason: "fixed-length extension content must be exactly 3 bytes",
});
}
out[1..WORD].copy_from_slice(self.content);
Ok(WORD)
} else {
if total % WORD != 0 {
return Err(Error::InvalidExtension {
reason: "variable-length extension total must be a multiple of 4 bytes",
});
}
let hel = total / WORD;
if hel > u8::MAX as usize {
return Err(Error::FieldTooWide {
what: "HEL",
value: hel as u64,
bits: 8,
});
}
out[1] = hel as u8;
out[2..total].copy_from_slice(self.content);
Ok(total)
}
}
}
pub fn parse_chain(mut data: &[u8]) -> Result<Vec<HeaderExtension<'_>>> {
let mut out = Vec::new();
while !data.is_empty() {
let (ext, n) = HeaderExtension::parse(data)?;
out.push(ext);
data = &data[n..];
}
Ok(out)
}
pub fn chain_len(exts: &[HeaderExtension<'_>]) -> usize {
exts.iter().map(|e| e.serialized_len()).sum()
}
pub fn serialize_chain(exts: &[HeaderExtension<'_>], out: &mut [u8]) -> Result<usize> {
let mut off = 0;
for e in exts {
off += e.serialize_into(&mut out[off..])?;
}
Ok(off)
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
#[test]
fn variable_ext_round_trip() {
let content = [0x11u8, 0x22, 0x33, 0x44, 0x55, 0x66];
let ext = HeaderExtension::new(0, &content);
assert!(!ext.is_fixed());
assert_eq!(ext.serialized_len(), 8);
assert_eq!(ext.hel(), 2);
let mut out = vec![0u8; ext.serialized_len()];
let n = ext.serialize_into(&mut out).unwrap();
assert_eq!(n, 8);
assert_eq!(&out, &[0x00, 0x02, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66]);
let (re, used) = HeaderExtension::parse(&out).unwrap();
assert_eq!(used, 8);
assert_eq!(re, ext);
}
#[test]
fn fixed_ext_round_trip() {
let content = [0xAAu8, 0xBB, 0xCC];
let ext = HeaderExtension::new(192, &content);
assert!(ext.is_fixed());
assert_eq!(ext.serialized_len(), 4);
let mut out = vec![0u8; 4];
ext.serialize_into(&mut out).unwrap();
assert_eq!(&out, &[0xC0, 0xAA, 0xBB, 0xCC]);
let (re, used) = HeaderExtension::parse(&out).unwrap();
assert_eq!(used, 4);
assert_eq!(re, ext);
}
#[test]
fn rejects_zero_hel() {
let data = [0x00u8, 0x00, 0x00, 0x00];
assert!(matches!(
HeaderExtension::parse(&data),
Err(Error::InvalidExtension { .. })
));
}
#[test]
fn multi_extension_chain_round_trips() {
let c1 = [0xDEu8, 0xAD, 0xBE, 0xEF, 0x00, 0x01];
let c2 = [0x02u8, 0x00, 0x00];
let exts = vec![HeaderExtension::new(1, &c1), HeaderExtension::new(128, &c2)];
let total = chain_len(&exts);
assert_eq!(total, 8 + 4);
let mut out = vec![0u8; total];
let n = serialize_chain(&exts, &mut out).unwrap();
assert_eq!(n, total);
let parsed = parse_chain(&out).unwrap();
assert_eq!(parsed, exts);
}
}