1use alloc::vec::Vec;
21
22use crate::error::{Error, Result};
23
24pub const FIXED_HET_MIN: u8 = 128;
27pub const WORD: usize = 4;
29
30#[derive(Debug, Clone, PartialEq, Eq)]
43#[cfg_attr(feature = "serde", derive(serde::Serialize))]
44pub struct HeaderExtension<'a> {
45 pub het: u8,
47 pub content: &'a [u8],
49}
50
51impl<'a> HeaderExtension<'a> {
52 pub fn new(het: u8, content: &'a [u8]) -> Self {
54 HeaderExtension { het, content }
55 }
56
57 pub fn is_fixed(&self) -> bool {
59 self.het >= FIXED_HET_MIN
60 }
61
62 pub fn serialized_len(&self) -> usize {
64 if self.is_fixed() {
65 WORD
66 } else {
67 2 + self.content.len()
69 }
70 }
71
72 pub fn hel(&self) -> usize {
75 self.serialized_len() / WORD
76 }
77
78 pub fn parse(data: &'a [u8]) -> Result<(Self, usize)> {
81 if data.is_empty() {
82 return Err(Error::BufferTooShort {
83 need: 1,
84 have: 0,
85 what: "header extension HET",
86 });
87 }
88 let het = data[0];
89 if het >= FIXED_HET_MIN {
90 if data.len() < WORD {
92 return Err(Error::BufferTooShort {
93 need: WORD,
94 have: data.len(),
95 what: "fixed-length header extension",
96 });
97 }
98 Ok((
99 HeaderExtension {
100 het,
101 content: &data[1..WORD],
102 },
103 WORD,
104 ))
105 } else {
106 if data.len() < 2 {
108 return Err(Error::BufferTooShort {
109 need: 2,
110 have: data.len(),
111 what: "variable-length header extension HEL",
112 });
113 }
114 let hel = data[1] as usize;
115 if hel == 0 {
116 return Err(Error::InvalidExtension {
117 reason: "HEL must be >= 1 for a variable-length extension",
118 });
119 }
120 let total = hel * WORD;
121 if data.len() < total {
122 return Err(Error::BufferTooShort {
123 need: total,
124 have: data.len(),
125 what: "variable-length header extension content",
126 });
127 }
128 Ok((
129 HeaderExtension {
130 het,
131 content: &data[2..total],
132 },
133 total,
134 ))
135 }
136 }
137
138 pub fn serialize_into(&self, out: &mut [u8]) -> Result<usize> {
141 let total = self.serialized_len();
142 if out.len() < total {
143 return Err(Error::OutputBufferTooSmall {
144 need: total,
145 have: out.len(),
146 });
147 }
148 out[0] = self.het;
149 if self.is_fixed() {
150 if self.content.len() != WORD - 1 {
151 return Err(Error::InvalidExtension {
152 reason: "fixed-length extension content must be exactly 3 bytes",
153 });
154 }
155 out[1..WORD].copy_from_slice(self.content);
156 Ok(WORD)
157 } else {
158 if !total.is_multiple_of(WORD) {
160 return Err(Error::InvalidExtension {
161 reason: "variable-length extension total must be a multiple of 4 bytes",
162 });
163 }
164 let hel = total / WORD;
165 if hel > u8::MAX as usize {
166 return Err(Error::FieldTooWide {
167 what: "HEL",
168 value: hel as u64,
169 bits: 8,
170 });
171 }
172 out[1] = hel as u8;
173 out[2..total].copy_from_slice(self.content);
174 Ok(total)
175 }
176 }
177}
178
179pub fn parse_chain(mut data: &[u8]) -> Result<Vec<HeaderExtension<'_>>> {
182 let mut out = Vec::new();
183 while !data.is_empty() {
184 let (ext, n) = HeaderExtension::parse(data)?;
185 out.push(ext);
186 data = &data[n..];
187 }
188 Ok(out)
189}
190
191pub fn chain_len(exts: &[HeaderExtension<'_>]) -> usize {
193 exts.iter().map(|e| e.serialized_len()).sum()
194}
195
196pub fn serialize_chain(exts: &[HeaderExtension<'_>], out: &mut [u8]) -> Result<usize> {
198 let mut off = 0;
199 for e in exts {
200 off += e.serialize_into(&mut out[off..])?;
201 }
202 Ok(off)
203}
204
205#[cfg(test)]
206mod tests {
207 use super::*;
208 use alloc::vec;
209
210 #[test]
211 fn variable_ext_round_trip() {
212 let content = [0x11u8, 0x22, 0x33, 0x44, 0x55, 0x66];
214 let ext = HeaderExtension::new(0, &content);
215 assert!(!ext.is_fixed());
216 assert_eq!(ext.serialized_len(), 8);
217 assert_eq!(ext.hel(), 2);
218
219 let mut out = vec![0u8; ext.serialized_len()];
220 let n = ext.serialize_into(&mut out).unwrap();
221 assert_eq!(n, 8);
222 assert_eq!(&out, &[0x00, 0x02, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66]);
223
224 let (re, used) = HeaderExtension::parse(&out).unwrap();
225 assert_eq!(used, 8);
226 assert_eq!(re, ext);
227 }
228
229 #[test]
230 fn fixed_ext_round_trip() {
231 let content = [0xAAu8, 0xBB, 0xCC];
233 let ext = HeaderExtension::new(192, &content);
234 assert!(ext.is_fixed());
235 assert_eq!(ext.serialized_len(), 4);
236
237 let mut out = vec![0u8; 4];
238 ext.serialize_into(&mut out).unwrap();
239 assert_eq!(&out, &[0xC0, 0xAA, 0xBB, 0xCC]);
240
241 let (re, used) = HeaderExtension::parse(&out).unwrap();
242 assert_eq!(used, 4);
243 assert_eq!(re, ext);
244 }
245
246 #[test]
247 fn rejects_zero_hel() {
248 let data = [0x00u8, 0x00, 0x00, 0x00];
249 assert!(matches!(
250 HeaderExtension::parse(&data),
251 Err(Error::InvalidExtension { .. })
252 ));
253 }
254
255 #[test]
256 fn multi_extension_chain_round_trips() {
257 let c1 = [0xDEu8, 0xAD, 0xBE, 0xEF, 0x00, 0x01];
259 let c2 = [0x02u8, 0x00, 0x00];
260 let exts = vec![HeaderExtension::new(1, &c1), HeaderExtension::new(128, &c2)];
261 let total = chain_len(&exts);
262 assert_eq!(total, 8 + 4);
263
264 let mut out = vec![0u8; total];
265 let n = serialize_chain(&exts, &mut out).unwrap();
266 assert_eq!(n, total);
267
268 let parsed = parse_chain(&out).unwrap();
269 assert_eq!(parsed, exts);
270 }
271}