use super::HpackError;
#[inline]
pub(crate) fn encode(out: &mut Vec<u8>, value: u64, prefix_bits: u8, header: u8) {
let mask = (1u64 << prefix_bits) - 1;
if value < mask {
out.push((u64::from(header) & !mask | value) as u8);
return;
}
out.push((u64::from(header) & !mask | mask) as u8);
let mut v = value - mask;
while v >= 128 {
out.push((v & 0x7f) as u8 | 0x80);
v >>= 7;
}
out.push(v as u8);
}
#[inline]
pub(crate) fn decode(
buf: &[u8],
off: &mut usize,
prefix_bits: u8,
header: u8,
) -> Result<u64, HpackError> {
let mask = (1u64 << prefix_bits) - 1;
let mut value = u64::from(header) & mask;
if value < mask {
return Ok(value);
}
let mut shift: u32 = 0;
loop {
if shift > 56 {
return Err(HpackError::InvalidInteger);
}
let b = *buf.get(*off).ok_or(HpackError::InvalidInteger)?;
*off += 1;
value = value
.checked_add((u64::from(b & 0x7f)) << shift)
.ok_or(HpackError::InvalidInteger)?;
if b & 0x80 == 0 {
if value >= (1 << 62) {
return Err(HpackError::InvalidInteger);
}
return Ok(value);
}
shift += 7;
}
}
#[allow(dead_code)]
#[inline]
pub(crate) fn encoded_len(buf: &[u8], prefix_bits: u8) -> Option<usize> {
let mask = (1u64 << prefix_bits) - 1;
let first = *buf.first()?;
if (u64::from(first) & mask) < mask {
return Some(1);
}
let mut len = 1;
loop {
let b = *buf.get(len)?;
len += 1;
if b & 0x80 == 0 {
return Some(len);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn encodes_with_prefix() {
let mut out = Vec::new();
encode(&mut out, 10, 5, 0);
assert_eq!(out, [0x0a]);
out.clear();
encode(&mut out, 1337, 5, 0);
assert_eq!(out, [0x1f, 0x9a, 0x0a]);
out.clear();
encode(&mut out, 42, 8, 0);
assert_eq!(out, [0x2a]);
}
#[test]
fn preserves_header_bits() {
let mut out = Vec::new();
encode(&mut out, 5, 7, 0x80);
assert_eq!(out, [0x85]);
out.clear();
encode(&mut out, 0, 7, 0x80);
assert_eq!(out, [0x80]);
out.clear();
encode(&mut out, 127, 7, 0x80);
assert_eq!(out, [0xff, 0x00]);
}
#[test]
fn round_trip() {
for &(bits, value) in &[
(5, 0u64),
(5, 30),
(5, 31),
(5, 1337),
(5, (1 << 62) - 1),
(6, 62),
(6, 63),
(6, 10_000),
(7, 127),
(7, 128),
(7, 1 << 24),
(7, 1 << 40),
(7, (1 << 62) - 1),
(8, 255),
(8, 256),
] {
let mut out = Vec::new();
encode(&mut out, value, bits, 0);
let mut off = 1;
assert_eq!(decode(&out, &mut off, bits, out[0]).unwrap(), value);
assert_eq!(off, out.len(), "value {value}");
}
}
#[test]
fn rejects_values_beyond_62_bits() {
let mut out = Vec::new();
encode(&mut out, u64::MAX, 5, 0);
let mut off = 1;
assert_eq!(
decode(&out, &mut off, 5, out[0]),
Err(HpackError::InvalidInteger)
);
out.clear();
encode(&mut out, 1 << 62, 5, 0);
let mut off = 1;
assert_eq!(
decode(&out, &mut off, 5, out[0]),
Err(HpackError::InvalidInteger)
);
}
#[test]
fn round_trip_qpack_prefix_sizes() {
for prefix_bits in 2..=8u8 {
let mut out = Vec::new();
encode(
&mut out,
2u64.pow(u32::from(prefix_bits)) + 123,
prefix_bits,
0,
);
let mut off = 1;
assert_eq!(
decode(&out, &mut off, prefix_bits, out[0]).unwrap(),
2u64.pow(u32::from(prefix_bits)) + 123
);
assert_eq!(off, out.len(), "prefix_bits {prefix_bits}");
}
}
#[test]
fn rejects_truncated() {
let mut out = Vec::new();
encode(&mut out, 1337, 5, 0);
let mut off = 1;
let header = 0x1f;
assert!(decode(&out[..1], &mut off, 5, header).is_err());
}
#[test]
fn rejects_overflow() {
let buf = [
0x1f, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0x7f,
];
let mut off = 1;
assert_eq!(
decode(&buf, &mut off, 5, buf[0]),
Err(HpackError::InvalidInteger)
);
}
}