use std::num::Wrapping;
use bytes::Bytes;
use super::HeaderTable;
use crate::hpack::static_table::StaticTable;
use crate::hpack::HeaderValueFound;
use bytes::BytesMut;
pub trait EncodeBuf {
fn write_all(&mut self, bytes: &[u8]);
fn reserve(&mut self, additional: usize) {
drop(additional);
}
fn write_u8(&mut self, b: u8) {
self.write_all(&[b]);
}
}
impl EncodeBuf for Vec<u8> {
fn write_all(&mut self, bytes: &[u8]) {
self.extend_from_slice(bytes);
}
fn reserve(&mut self, additional: usize) {
self.reserve(additional);
}
fn write_u8(&mut self, byte: u8) {
self.push(byte);
}
}
impl EncodeBuf for BytesMut {
fn write_all(&mut self, bytes: &[u8]) {
self.extend_from_slice(bytes);
}
fn reserve(&mut self, additional: usize) {
self.reserve(additional);
}
}
pub fn encode_integer_into<W: EncodeBuf>(
mut value: usize,
prefix_size: u8,
leading_bits: u8,
writer: &mut W,
) {
let Wrapping(mask) = if prefix_size >= 8 {
Wrapping(0xFF)
} else {
Wrapping(1u8 << prefix_size) - Wrapping(1)
};
let leading_bits = leading_bits & (!mask);
let mask = mask as usize;
if value < mask {
writer.write_u8(leading_bits | value as u8);
return;
}
writer.write_u8(leading_bits | mask as u8);
value -= mask;
while value >= 128 {
writer.write_u8(((value % 128) + 128) as u8);
value = value / 128;
}
writer.write_u8(value as u8);
}
#[cfg(test)]
pub fn encode_integer(value: usize, prefix_size: u8) -> Vec<u8> {
let mut res = Vec::new();
encode_integer_into(value, prefix_size, 0, &mut res);
res
}
pub struct Encoder {
header_table: HeaderTable,
}
impl Encoder {
pub fn new() -> Encoder {
Encoder {
header_table: HeaderTable::with_static_table(StaticTable::new()),
}
}
pub fn encode_for_test<'b, I>(&mut self, headers: I) -> Vec<u8>
where
I: IntoIterator<Item = (&'b [u8], &'b [u8])>,
{
let mut encoded: Vec<u8> = Vec::new();
self.encode_into(headers, &mut encoded);
encoded
}
pub fn encode<'b, I>(&mut self, headers: I) -> Bytes
where
I: IntoIterator<Item = (&'b [u8], &'b [u8])>,
{
let mut encoded = BytesMut::new();
self.encode_into(headers, &mut encoded);
encoded.freeze()
}
pub fn encode_into<'b, I, W>(&mut self, headers: I, writer: &mut W)
where
I: IntoIterator<Item = (&'b [u8], &'b [u8])>,
W: EncodeBuf,
{
for header in headers {
self.encode_header_into(header, writer);
}
}
fn encode_header_into<W: EncodeBuf>(&mut self, header: (&[u8], &[u8]), writer: &mut W) {
match self.header_table.find_header(header) {
None => {
self.encode_literal(&header, true, writer);
self.header_table.add_header(
Bytes::copy_from_slice(header.0),
Bytes::copy_from_slice(header.1),
);
}
Some((index, HeaderValueFound::NameOnlyFound)) => {
self.encode_indexed_name((index, header.1), false, writer);
}
Some((index, HeaderValueFound::Found)) => {
self.encode_indexed(index, writer);
}
};
}
fn encode_literal<W: EncodeBuf>(
&mut self,
header: &(&[u8], &[u8]),
should_index: bool,
buf: &mut W,
) {
let mask = if should_index { 0x40 } else { 0x0 };
buf.write_u8(mask);
self.encode_string_literal(&header.0, buf);
self.encode_string_literal(&header.1, buf);
}
fn encode_string_literal<W: EncodeBuf>(&mut self, octet_str: &[u8], buf: &mut W) {
buf.reserve(octet_str.len() + 1);
encode_integer_into(octet_str.len(), 7, 0, buf);
buf.write_all(octet_str);
}
fn encode_indexed_name<W: EncodeBuf>(
&mut self,
header: (usize, &[u8]),
should_index: bool,
buf: &mut W,
) {
let (mask, prefix) = if should_index { (0x40, 6) } else { (0x0, 4) };
encode_integer_into(header.0, prefix, mask, buf);
self.encode_string_literal(&header.1, buf);
}
fn encode_indexed<W: EncodeBuf>(&self, index: usize, buf: &mut W) {
encode_integer_into(index, 7, 0x80, buf);
}
}
#[cfg(test)]
mod tests {
use bytes::Bytes;
use super::encode_integer;
use super::Encoder;
use super::super::Decoder;
#[test]
fn test_encode_integer() {
assert_eq!(encode_integer(10, 5), [10]);
assert_eq!(encode_integer(1337, 5), [31, 154, 10]);
assert_eq!(encode_integer(127, 7), [127, 0]);
assert_eq!(encode_integer(255, 8), [255, 0]);
assert_eq!(encode_integer(254, 8), [254]);
assert_eq!(encode_integer(1, 8), [1]);
assert_eq!(encode_integer(0, 8), [0]);
assert_eq!(encode_integer(255, 7), [127, 128, 1]);
}
fn is_decodable(buf: &Vec<u8>, headers: &[(Vec<u8>, Vec<u8>)]) -> bool {
let mut decoder = Decoder::new();
match decoder.decode_for_test(&buf[..]).ok() {
Some(h) => {
h == headers
.iter()
.map(|(k, v)| {
(
Bytes::copy_from_slice(&(*k)[..]),
Bytes::copy_from_slice(&(*v)[..]),
)
})
.collect::<Vec<_>>()
}
None => false,
}
}
#[test]
fn test_encode_only_method() {
let mut encoder: Encoder = Encoder::new();
let headers = vec![(b":method".to_vec(), b"GET".to_vec())];
let result = encoder.encode_for_test(headers.iter().map(|h| (&h.0[..], &h.1[..])));
debug!("{:?}", result);
assert!(is_decodable(&result, &headers));
}
#[test]
fn test_custom_header_gets_indexed() {
let mut encoder: Encoder = Encoder::new();
let headers = vec![(b"custom-key".to_vec(), b"custom-value".to_vec())];
let result = encoder.encode_for_test(headers.iter().map(|h| (&h.0[..], &h.1[..])));
assert!(is_decodable(&result, &headers));
assert_eq!(encoder.header_table.dynamic_table.to_vec_of_vec(), headers);
assert!(0x40 == (0x40 & result[0]));
debug!("{:?}", result);
}
#[test]
fn test_uses_index_on_second_iteration() {
let mut encoder: Encoder = Encoder::new();
let headers = vec![(b"custom-key".to_vec(), b"custom-value".to_vec())];
let _ = encoder.encode_for_test(headers.iter().map(|h| (&h.0[..], &h.1[..])));
let result = encoder.encode_for_test(headers.iter().map(|h| (&h.0[..], &h.1[..])));
assert_eq!(encoder.header_table.dynamic_table.to_vec_of_vec(), headers);
assert_eq!(result.len(), 1);
assert_eq!(0x80 & result[0], 0x80);
assert_eq!(result[0] ^ 0x80, 62);
assert_eq!(
encoder.header_table.get_from_table_vec(62).unwrap(),
headers[0]
);
}
#[test]
fn test_name_indexed_value_not() {
{
let mut encoder: Encoder = Encoder::new();
let headers = vec![(b":method", b"PUT")];
let result = encoder.encode_for_test(headers.iter().map(|h| (&h.0[..], &h.1[..])));
assert_eq!(result[0], 3);
assert_eq!(&result[1..], &[3, b'P', b'U', b'T']);
}
{
let mut encoder: Encoder = Encoder::new();
let headers = vec![(b":authority".to_vec(), b"example.com".to_vec())];
let result = encoder.encode_for_test(headers.iter().map(|h| (&h.0[..], &h.1[..])));
assert_eq!(result[0], 1);
assert_eq!(
&result[1..],
&[11, b'e', b'x', b'a', b'm', b'p', b'l', b'e', b'.', b'c', b'o', b'm']
)
}
}
#[test]
fn test_multiple_headers_encoded() {
let mut encoder = Encoder::new();
let headers = vec![
(b"custom-key".to_vec(), b"custom-value".to_vec()),
(b":method".to_vec(), b"GET".to_vec()),
(b":path".to_vec(), b"/some/path".to_vec()),
];
let result = encoder.encode_for_test(headers.iter().map(|h| (&h.0[..], &h.1[..])));
assert!(is_decodable(&result, &headers));
}
}