pub mod huffman;
pub mod huffman_table;
pub mod table;
use crate::courierust_bytes::{Bytes, BytesMut};
use crate::courierust_error::{Error, Result};
use crate::courierust_hpack::huffman::{encode as huffman_encode, HuffmanDecoder};
use crate::courierust_hpack::table::Table;
use crate::courierust_http::header::{HeaderName, HeaderValue};
use alloc::vec::Vec;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct HeaderField {
pub name: HeaderName,
pub value: HeaderValue,
pub never_indexed: bool,
}
impl HeaderField {
pub fn new(name: HeaderName, value: HeaderValue) -> Self {
Self {
name,
value,
never_indexed: false,
}
}
pub fn new_never_indexed(name: HeaderName, value: HeaderValue) -> Self {
Self {
name,
value,
never_indexed: true,
}
}
}
pub type HeaderList = Vec<HeaderField>;
const SENSITIVE_NAMES: [&str; 4] = [
"authorization",
"cookie",
"proxy-authorization",
"set-cookie",
];
fn read_int(input: &[u8], pos: &mut usize, n: u8, prefix: u8) -> Result<usize> {
let max_prefix = (1u16 << n) - 1;
let mut val = prefix as usize;
if (prefix as u16) < max_prefix {
return Ok(val);
}
let mut shift = 0usize;
loop {
if *pos >= input.len() {
return Err(Error::protocol("HPACK: truncated integer"));
}
let b = input[*pos];
*pos += 1;
val = val
.checked_add(((b & 0x7f) as usize) << shift)
.ok_or_else(|| Error::overflow("HPACK: integer overflow"))?;
if b & 0x80 == 0 {
return Ok(val);
}
shift += 7;
if shift > 63 {
return Err(Error::overflow("HPACK: integer too large"));
}
}
}
fn read_string(
input: &[u8],
pos: &mut usize,
max_len: usize,
huff: &HuffmanDecoder,
) -> Result<Bytes> {
if *pos >= input.len() {
return Err(Error::protocol("HPACK: truncated string"));
}
let b = input[*pos];
*pos += 1;
let huffman = b & 0x80 != 0;
let len = read_int(input, pos, 7, b & 0x7f)?;
if len > max_len {
return Err(Error::overflow("HPACK: string exceeds limit"));
}
if *pos + len > input.len() {
return Err(Error::protocol("HPACK: truncated string data"));
}
let raw = &input[*pos..*pos + len];
*pos += len;
if huffman {
let mut out = Vec::with_capacity(len);
huff.decode(raw, &mut out, len.saturating_mul(2))
.map_err(|e| Error::protocol(format!("HPACK: Huffman decode error: {e:?}")))?;
Ok(Bytes::from(out))
} else {
Ok(Bytes::from(raw))
}
}
pub struct Decoder {
table: Table,
huff: HuffmanDecoder,
max_table_size: usize,
max_header_list_size: usize,
max_string_size: usize,
}
impl Decoder {
pub fn new(max_table_size: usize, max_header_list_size: usize) -> Self {
Self {
table: Table::default(),
huff: HuffmanDecoder::new(),
max_table_size,
max_header_list_size,
max_string_size: core::cmp::max(max_header_list_size.saturating_mul(4), 1 << 20),
}
}
#[inline]
pub fn table_size(&self) -> usize {
self.table.size()
}
pub fn decode(&mut self, input: &[u8]) -> Result<HeaderList> {
let mut out = HeaderList::new();
let mut total = 0usize;
let mut pos = 0usize;
let mut saw_rep = false;
while pos < input.len() {
let b = input[pos];
pos += 1;
if b & 0x80 != 0 {
let idx = read_int(input, &mut pos, 7, b & 0x7f)?;
if idx == 0 {
return Err(Error::protocol("HPACK: indexed with index 0"));
}
let (n, v) = self
.table
.get(idx)
.ok_or_else(|| Error::protocol("HPACK: index out of range"))?;
let name = HeaderName::from_hpack_bytes(n)?;
let value = HeaderValue::from_bytes(v)?;
total = checked_add(total, n.len() + v.len())?;
if total > self.max_header_list_size {
return Err(Error::overflow("HPACK: header list too large"));
}
out.push(HeaderField::new(name, value));
saw_rep = true;
} else if b & 0x40 != 0 {
let name_idx = read_int(input, &mut pos, 6, b & 0x3f)?;
let (name_bytes, name_len) = if name_idx == 0 {
let s = read_string(input, &mut pos, self.max_string_size, &self.huff)?;
let l = s.len();
(s, l)
} else {
let (n, _) = self
.table
.get(name_idx)
.ok_or_else(|| Error::protocol("HPACK: name index out of range"))?;
(Bytes::from(n), n.len())
};
let name = HeaderName::from_hpack_bytes(name_bytes.as_slice())?;
let value = read_string(input, &mut pos, self.max_string_size, &self.huff)?;
let vbytes = value.as_slice();
let value = HeaderValue::from_bytes(vbytes)?;
total = checked_add(total, name_len + vbytes.len())?;
if total > self.max_header_list_size {
return Err(Error::overflow("HPACK: header list too large"));
}
self.table.dynamic().insert(name.as_bytes(), vbytes);
out.push(HeaderField::new(name, value));
saw_rep = true;
} else if b & 0x20 != 0 {
if saw_rep {
return Err(Error::protocol(
"HPACK: size update after field representations",
));
}
let new = read_int(input, &mut pos, 5, b & 0x1f)?;
if new > self.max_table_size {
return Err(Error::protocol(
"HPACK: size update exceeds advertised maximum",
));
}
self.table.dynamic().set_max_size(new);
} else {
let never_indexed = b & 0x10 != 0;
let name_idx = read_int(input, &mut pos, 4, b & 0x0f)?;
let (name_bytes, name_len) = if name_idx == 0 {
let s = read_string(input, &mut pos, self.max_string_size, &self.huff)?;
let l = s.len();
(s, l)
} else {
let (n, _) = self
.table
.get(name_idx)
.ok_or_else(|| Error::protocol("HPACK: name index out of range"))?;
(Bytes::from(n), n.len())
};
let name = HeaderName::from_hpack_bytes(name_bytes.as_slice())?;
let value = read_string(input, &mut pos, self.max_string_size, &self.huff)?;
let vbytes = value.as_slice();
let value = HeaderValue::from_bytes(vbytes)?;
total = checked_add(total, name_len + vbytes.len())?;
if total > self.max_header_list_size {
return Err(Error::overflow("HPACK: header list too large"));
}
out.push(HeaderField {
name,
value,
never_indexed,
});
saw_rep = true;
}
}
Ok(out)
}
}
#[inline]
fn checked_add(a: usize, b: usize) -> Result<usize> {
a.checked_add(b)
.ok_or_else(|| Error::overflow("HPACK: size arithmetic overflow"))
}
pub struct Encoder {
table: Table,
peer_max_table_size: usize,
pending_size_update: Option<usize>,
max_index_len: usize,
}
impl Default for Encoder {
fn default() -> Self {
Self::new()
}
}
impl Encoder {
pub fn new() -> Self {
Self {
table: Table::default(),
peer_max_table_size: 4096,
pending_size_update: None,
max_index_len: 256,
}
}
pub fn set_peer_table_size(&mut self, size: usize) {
self.peer_max_table_size = size;
self.pending_size_update = Some(size);
}
#[inline]
pub fn table_size(&self) -> usize {
self.table.size()
}
pub fn encode(&mut self, fields: &[HeaderField], out: &mut BytesMut) {
if let Some(new) = self.pending_size_update.take() {
write_table_size_update(new, out);
self.table.dynamic().set_max_size(new);
}
for f in fields {
let name = f.name.as_bytes();
let value = f.value.as_bytes();
let sensitive = f.never_indexed || is_sensitive_name(name);
if !sensitive {
if let Some(idx) = self.table.find_full(name, value) {
write_int(0x80, 7, idx, out);
continue;
}
}
let name_idx = self.table.find_name(name);
let want_index = !sensitive && should_index(value, self.max_index_len);
if sensitive {
write_literal(0x10, 4, name_idx, name, value, out);
} else if want_index {
write_literal(0x40, 6, name_idx, name, value, out);
self.table.dynamic().insert(name, value);
} else {
write_literal(0x00, 4, name_idx, name, value, out);
}
}
}
}
#[inline]
fn is_sensitive_name(name: &[u8]) -> bool {
SENSITIVE_NAMES.iter().any(|s| s.as_bytes() == name)
}
#[inline]
fn should_index(value: &[u8], max_index_len: usize) -> bool {
!value.is_empty() && value.len() <= max_index_len
}
fn write_int(prefix: u8, n: u8, value: usize, out: &mut BytesMut) {
let max_prefix = (1u16 << n) - 1;
if (value as u16) < max_prefix {
out.put_u8(prefix | value as u8);
return;
}
out.put_u8(prefix | max_prefix as u8);
let mut v = (value as u64) - max_prefix as u64;
while v >= 128 {
out.put_u8(((v % 128) as u8) | 0x80);
v /= 128;
}
out.put_u8(v as u8);
}
fn write_string(s: &[u8], out: &mut BytesMut) {
let mut huff = Vec::with_capacity(s.len());
huffman_encode(s, &mut huff);
if huff.len() < s.len() {
write_int(0x80, 7, huff.len(), out);
out.extend_from_slice(&huff);
} else {
write_int(0x00, 7, s.len(), out);
out.extend_from_slice(s);
}
}
fn write_literal(
prefix: u8,
n: u8,
name_idx: Option<usize>,
name: &[u8],
value: &[u8],
out: &mut BytesMut,
) {
match name_idx {
Some(idx) => write_int(prefix, n, idx, out),
None => {
write_int(prefix, n, 0, out);
write_string(name, out);
}
}
write_string(value, out);
}
fn write_table_size_update(size: usize, out: &mut BytesMut) {
write_int(0x20, 5, size, out);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::courierust_error::ErrorKind;
use crate::courierust_hpack::huffman_table::HUFFMAN;
use alloc::string::String;
#[test]
fn read_int_overflow_rejected() {
let input = [0xffu8; 10];
let mut pos = 0usize;
let r = read_int(&input, &mut pos, 7, 0x7f);
assert!(r.is_err());
assert!(matches!(
r,
Err(Error {
kind: ErrorKind::Overflow,
..
})
));
}
fn hex(s: &str) -> Vec<u8> {
let s: String = s.chars().filter(|c| !c.is_whitespace()).collect();
assert!(s.len().is_multiple_of(2));
(0..s.len() / 2)
.map(|i| u8::from_str_radix(&s[i * 2..i * 2 + 2], 16).unwrap())
.collect()
}
fn headers(list: &[(&str, &str)]) -> HeaderList {
list.iter()
.map(|(n, v)| {
HeaderField::new(
HeaderName::from_hpack_bytes(n.as_bytes()).unwrap(),
HeaderValue::from_bytes(v.as_bytes()).unwrap(),
)
})
.collect()
}
#[test]
fn rfc_c3_request_sequence_no_huffman() {
let mut dec = Decoder::new(4096, 1 << 20);
let out1 = dec
.decode(&hex("8286 8441 0f77 7777 2e65 7861 6d70 6c65 2e63 6f6d"))
.unwrap();
assert_eq!(
out1,
headers(&[
(":method", "GET"),
(":scheme", "http"),
(":path", "/"),
(":authority", "www.example.com"),
])
);
let out2 = dec
.decode(&hex("8286 84be 5808 6e6f 2d63 6163 6865"))
.unwrap();
assert_eq!(
out2,
headers(&[
(":method", "GET"),
(":scheme", "http"),
(":path", "/"),
(":authority", "www.example.com"),
("cache-control", "no-cache"),
])
);
let out3 = dec
.decode(&hex(
"8287 85bf 400a 6375 7374 6f6d 2d6b 6579 0c63 7573 746f 6d2d 7661 6c75 65",
))
.unwrap();
assert_eq!(
out3,
headers(&[
(":method", "GET"),
(":scheme", "https"),
(":path", "/index.html"),
(":authority", "www.example.com"),
("custom-key", "custom-value"),
])
);
}
#[test]
fn rfc_c4_request_sequence_huffman() {
let mut dec = Decoder::new(4096, 1 << 20);
let out1 = dec
.decode(&hex("8286 8441 8cf1 e3c2 e5f2 3a6b a0ab 90f4 ff"))
.unwrap();
assert_eq!(
out1,
headers(&[
(":method", "GET"),
(":scheme", "http"),
(":path", "/"),
(":authority", "www.example.com"),
])
);
let out2 = dec.decode(&hex("8286 84be 5886 a8eb 1064 9cbf")).unwrap();
assert_eq!(
out2,
headers(&[
(":method", "GET"),
(":scheme", "http"),
(":path", "/"),
(":authority", "www.example.com"),
("cache-control", "no-cache"),
])
);
let out3 = dec
.decode(&hex(
"8287 85bf 4088 25a8 49e9 5ba9 7d7f 8925 a849 e95b b8e8 b4bf",
))
.unwrap();
assert_eq!(
out3,
headers(&[
(":method", "GET"),
(":scheme", "https"),
(":path", "/index.html"),
(":authority", "www.example.com"),
("custom-key", "custom-value"),
])
);
}
#[test]
fn rfc_c6_response_sequence_huffman() {
let mut dec = Decoder::new(256, 1 << 20);
let out1 = dec
.decode(&hex("4882 6402 5885 aec3 771a 4b61 96d0 7abe \
9410 54d4 44a8 2005 9504 0b81 66e0 82a6 \
2d1b ff6e 919d 29ad 1718 63c7 8f0b 97c8 \
e9ae 82ae 43d3"))
.unwrap();
assert_eq!(
out1,
headers(&[
(":status", "302"),
("cache-control", "private"),
("date", "Mon, 21 Oct 2013 20:13:21 GMT"),
("location", "https://www.example.com"),
])
);
let out2 = dec.decode(&hex("4883 640e ffc1 c0bf")).unwrap();
assert_eq!(
out2,
headers(&[
(":status", "307"),
("cache-control", "private"),
("date", "Mon, 21 Oct 2013 20:13:21 GMT"),
("location", "https://www.example.com"),
])
);
let out3 = dec
.decode(&hex("88c1 6196 d07a be94 1054 d444 a820 0595 \
040b 8166 e084 a62d 1bff c05a 839b d9ab \
77ad 94e7 821d d7f2 e6c7 b335 dfdf cd5b \
3960 d5af 2708 7f36 72c1 ab27 0fb5 291f \
9587 3160 65c0 03ed 4ee5 b106 3d50 07"))
.unwrap();
assert_eq!(
out3,
headers(&[
(":status", "200"),
("cache-control", "private"),
("date", "Mon, 21 Oct 2013 20:13:22 GMT"),
("location", "https://www.example.com"),
("content-encoding", "gzip"),
(
"set-cookie",
"foo=ASDJKHQKBZXOQWEOPIUAXQWEOIU; max-age=3600; version=1"
),
])
);
}
#[test]
fn rfc_c2_literals() {
let mut dec = Decoder::new(4096, 1 << 20);
let out = dec
.decode(&hex(
"400a 6375 7374 6f6d 2d6b 6579 0d63 7573 746f 6d2d 6865 6164 6572",
))
.unwrap();
assert_eq!(out, headers(&[("custom-key", "custom-header")]));
let out2 = dec
.decode(&hex("040c 2f73 616d 706c 652f 7061 7468"))
.unwrap();
assert_eq!(out2, headers(&[(":path", "/sample/path")]));
let out3 = dec
.decode(&hex("1008 7061 7373 776f 7264 0673 6563 7265 74"))
.unwrap();
assert_eq!(out3.len(), 1);
assert_eq!(out3[0].name.as_str(), "password");
assert_eq!(out3[0].value.as_bytes(), b"secret");
assert!(out3[0].never_indexed);
let out4 = dec.decode(&hex("82")).unwrap();
assert_eq!(out4, headers(&[(":method", "GET")]));
}
#[test]
fn huffman_table_is_complete() {
assert_eq!(HUFFMAN.len(), 257);
for (i, &(code, len)) in HUFFMAN.iter().enumerate() {
assert!((5..=30).contains(&(len as u16)), "bad len {len} at sym {i}");
assert!(
code < (1u32 << len),
"code does not fit len at sym {i}: code={code:#x} len={len}"
);
}
}
#[test]
fn encoder_decoder_roundtrip() {
let mut enc = Encoder::new();
let fields = headers(&[
(":method", "GET"),
(":scheme", "https"),
(":path", "/index.html"),
(":authority", "www.example.com"),
("accept-encoding", "gzip, deflate, br"),
("cache-control", "no-cache"),
]);
let mut wire = BytesMut::new();
enc.encode(&fields, &mut wire);
let mut dec = Decoder::new(4096, 1 << 20);
let back = dec.decode(wire.as_slice()).unwrap();
assert_eq!(back, fields);
}
#[test]
fn encoder_reuses_dynamic_table() {
let mut enc = Encoder::new();
let fields = headers(&[("x-foo", "bar"), ("x-foo", "bar")]);
let mut wire = BytesMut::new();
enc.encode(&fields, &mut wire);
let mut dec = Decoder::new(4096, 1 << 20);
let back = dec.decode(wire.as_slice()).unwrap();
assert_eq!(back, fields);
assert!(
wire.len() < 32,
"expected compact encoding, got {} bytes",
wire.len()
);
}
#[test]
fn never_indexed_survives() {
let mut enc = Encoder::new();
let fields = vec![HeaderField::new_never_indexed(
HeaderName::from_static("authorization"),
HeaderValue::from_static("Bearer secret-token"),
)];
let mut wire = BytesMut::new();
enc.encode(&fields, &mut wire);
let mut dec = Decoder::new(4096, 1 << 20);
let back = dec.decode(wire.as_slice()).unwrap();
assert_eq!(back, fields);
}
#[test]
fn rejects_index_zero_and_bad_update() {
let mut dec = Decoder::new(4096, 1 << 20);
assert!(dec.decode(&[0x80]).is_err());
let mut dec2 = Decoder::new(4096, 1 << 20);
let mut wire = BytesMut::new();
write_table_size_update(5000, &mut wire);
assert!(dec2.decode(wire.as_slice()).is_err());
}
}