1#[doc(hidden)]
12pub mod huffman;
13mod huffman_table;
14mod static_table;
15
16use std::collections::HashMap;
17use std::sync::OnceLock;
18
19use huffman::{huffman_decode, huffman_encode, huffman_shorter};
20use static_table::{STATIC_TABLE, STATIC_TABLE_LENGTH};
21
22use crate::bytes::{decode_utf8, ByteReader, ByteWriter};
23use crate::errors::{ErrorCode, H2Error};
24
25const ENTRY_OVERHEAD: usize = 32; const DEFAULT_TABLE_SIZE: usize = 4096;
27
28#[derive(Debug, Clone, PartialEq, Eq)]
30pub struct Header {
31 pub name: String,
32 pub value: String,
33 pub never_index: bool,
35}
36
37impl Header {
38 pub fn new(name: impl Into<String>, value: impl Into<String>) -> Self {
39 Self {
40 name: name.into(),
41 value: value.into(),
42 never_index: false,
43 }
44 }
45
46 pub fn never_indexed(name: impl Into<String>, value: impl Into<String>) -> Self {
47 Self {
48 name: name.into(),
49 value: value.into(),
50 never_index: true,
51 }
52 }
53}
54
55fn write_integer(w: &mut ByteWriter, value: usize, prefix_bits: u32, first_byte_flags: u8) {
59 let max = (1usize << prefix_bits) - 1;
60 if value < max {
61 w.u8(first_byte_flags | value as u8);
62 return;
63 }
64 w.u8(first_byte_flags | max as u8);
65 let mut rest = value - max;
66 while rest >= 128 {
67 w.u8((rest % 128) as u8 + 128);
68 rest /= 128;
69 }
70 w.u8(rest as u8);
71}
72
73fn read_integer(r: &mut ByteReader, prefix_bits: u32) -> Result<usize, H2Error> {
74 let max = (1usize << prefix_bits) - 1;
75 let mut value = (r.u8() & max as u8) as usize;
76 if value < max {
77 return Ok(value);
78 }
79 let mut shift = 0u32;
80 loop {
81 if r.remaining() == 0 {
82 return Err(H2Error::new(
83 ErrorCode::CompressionError,
84 "truncated integer",
85 ));
86 }
87 let byte = r.u8();
88 value += ((byte & 0x7f) as usize) << shift;
89 shift += 7;
90 if shift > 42 {
91 return Err(H2Error::new(
92 ErrorCode::CompressionError,
93 "integer overflow",
94 ));
95 }
96 if byte & 0x80 == 0 {
97 break;
98 }
99 }
100 Ok(value)
101}
102
103fn write_string(w: &mut ByteWriter, s: &str) {
106 let raw = s.as_bytes();
107 if huffman_shorter(raw) {
108 let encoded = huffman_encode(raw);
109 write_integer(w, encoded.len(), 7, 0x80);
110 w.bytes(&encoded);
111 } else {
112 write_integer(w, raw.len(), 7, 0x00);
113 w.bytes(raw);
114 }
115}
116
117fn read_string(r: &mut ByteReader) -> Result<String, H2Error> {
118 if r.remaining() == 0 {
119 return Err(H2Error::new(
120 ErrorCode::CompressionError,
121 "truncated string",
122 ));
123 }
124 let huffman = r.peek() & 0x80 != 0;
125 let length = read_integer(r, 7)?;
126 if length > r.remaining() {
127 return Err(H2Error::new(
128 ErrorCode::CompressionError,
129 "string length exceeds block",
130 ));
131 }
132 let raw = r
133 .bytes(length)
134 .ok_or_else(|| H2Error::new(ErrorCode::CompressionError, "truncated string"))?;
135 let decoded = if huffman {
136 huffman_decode(raw)?
137 } else {
138 raw.to_vec()
139 };
140 Ok(decode_utf8(&decoded))
141}
142
143struct StaticMaps {
146 name_to_index: HashMap<String, usize>,
148 pair_to_index: HashMap<String, usize>,
150}
151
152fn pair_key(name: &str, value: &str) -> String {
153 format!("{name}\u{0}{value}")
154}
155
156fn static_maps() -> &'static StaticMaps {
157 static MAPS: OnceLock<StaticMaps> = OnceLock::new();
158 MAPS.get_or_init(|| {
159 let mut name_to_index = HashMap::new();
160 let mut pair_to_index = HashMap::new();
161 for (i, &(name, value)) in STATIC_TABLE.iter().enumerate() {
162 let index = i + 1;
163 name_to_index.entry(name.to_string()).or_insert(index);
164 pair_to_index.insert(pair_key(name, value), index);
165 }
166 StaticMaps {
167 name_to_index,
168 pair_to_index,
169 }
170 })
171}
172
173pub struct HpackDecoder {
176 dynamic: Vec<(String, String)>, size: usize,
178 max_size: usize,
179 protocol_max: usize,
180}
181
182impl Default for HpackDecoder {
183 fn default() -> Self {
184 Self::new(DEFAULT_TABLE_SIZE)
185 }
186}
187
188impl HpackDecoder {
189 pub fn new(max_size: usize) -> Self {
190 Self {
191 dynamic: Vec::new(),
192 size: 0,
193 max_size,
194 protocol_max: max_size,
195 }
196 }
197
198 pub fn set_protocol_max_size(&mut self, n: usize) {
200 self.protocol_max = n;
201 if self.max_size > n {
202 let _ = self.apply_max_size(n);
204 }
205 }
206
207 fn entry_at(&self, index: usize) -> Result<(String, String), H2Error> {
208 if (1..=STATIC_TABLE_LENGTH).contains(&index) {
209 let (n, v) = STATIC_TABLE[index - 1];
210 return Ok((n.to_string(), v.to_string()));
211 }
212 let di = index - STATIC_TABLE_LENGTH - 1;
213 self.dynamic.get(di).cloned().ok_or_else(|| {
214 H2Error::new(
215 ErrorCode::CompressionError,
216 format!("invalid HPACK index {index}"),
217 )
218 })
219 }
220
221 fn insert(&mut self, name: String, value: String) {
222 let entry_size = name.len() + value.len() + ENTRY_OVERHEAD;
223 while self.size + entry_size > self.max_size && !self.dynamic.is_empty() {
224 let removed = self.dynamic.pop().unwrap();
225 self.size -= removed.0.len() + removed.1.len() + ENTRY_OVERHEAD;
226 }
227 if entry_size <= self.max_size {
228 self.dynamic.insert(0, (name, value));
229 self.size += entry_size;
230 } else {
231 self.dynamic.clear();
233 self.size = 0;
234 }
235 }
236
237 fn apply_max_size(&mut self, new_size: usize) -> Result<(), H2Error> {
238 if new_size > self.protocol_max {
239 return Err(H2Error::new(
240 ErrorCode::CompressionError,
241 "dynamic table size update too large",
242 ));
243 }
244 self.max_size = new_size;
245 while self.size > self.max_size && !self.dynamic.is_empty() {
246 let removed = self.dynamic.pop().unwrap();
247 self.size -= removed.0.len() + removed.1.len() + ENTRY_OVERHEAD;
248 }
249 Ok(())
250 }
251
252 pub fn decode(&mut self, block: &[u8]) -> Result<Vec<Header>, H2Error> {
254 let mut r = ByteReader::new(block);
255 let mut headers = Vec::new();
256
257 while r.remaining() > 0 {
258 let first = r.peek();
259 if first & 0x80 != 0 {
260 let index = read_integer(&mut r, 7)?;
262 if index == 0 {
263 return Err(H2Error::new(
264 ErrorCode::CompressionError,
265 "indexed field with index 0",
266 ));
267 }
268 let (name, value) = self.entry_at(index)?;
269 headers.push(Header {
270 name,
271 value,
272 never_index: false,
273 });
274 } else if first & 0x40 != 0 {
275 let name_index = read_integer(&mut r, 6)?;
277 let name = if name_index == 0 {
278 read_string(&mut r)?
279 } else {
280 self.entry_at(name_index)?.0
281 };
282 let value = read_string(&mut r)?;
283 self.insert(name.clone(), value.clone());
284 headers.push(Header {
285 name,
286 value,
287 never_index: false,
288 });
289 } else if first & 0x20 != 0 {
290 let new_size = read_integer(&mut r, 5)?;
292 self.apply_max_size(new_size)?;
293 } else {
294 let name_index = read_integer(&mut r, 4)?;
296 let name = if name_index == 0 {
297 read_string(&mut r)?
298 } else {
299 self.entry_at(name_index)?.0
300 };
301 let value = read_string(&mut r)?;
302 headers.push(Header {
303 name,
304 value,
305 never_index: false,
306 });
307 }
308 }
309 Ok(headers)
310 }
311}
312
313#[derive(Default)]
315pub struct HpackEncoder;
316
317impl HpackEncoder {
318 pub fn new() -> Self {
319 Self
320 }
321
322 pub fn encode(&self, headers: &[Header]) -> Vec<u8> {
323 let maps = static_maps();
324 let mut w = ByteWriter::with_capacity(256);
325 for header in headers {
326 let name = header.name.to_ascii_lowercase();
327 let value = &header.value;
328
329 if let Some(&index) = maps.pair_to_index.get(&pair_key(&name, value)) {
330 write_integer(&mut w, index, 7, 0x80);
332 continue;
333 }
334
335 let flags = if header.never_index { 0x10 } else { 0x00 };
338 if let Some(&name_index) = maps.name_to_index.get(&name) {
339 write_integer(&mut w, name_index, 4, flags);
340 } else {
341 write_integer(&mut w, 0, 4, flags);
342 write_string(&mut w, &name);
343 }
344 write_string(&mut w, value);
345 }
346 w.into_vec()
347 }
348}