use bytes::{Buf, BufMut, Bytes, BytesMut};
pub const MAX_VARINT: u64 = (1 << 62) - 1;
pub const FRAME_DATA: u64 = 0x0;
pub const FRAME_HEADERS: u64 = 0x1;
pub const FRAME_CANCEL_PUSH: u64 = 0x3;
pub const FRAME_SETTINGS: u64 = 0x4;
pub const FRAME_PUSH_PROMISE: u64 = 0x5;
pub const FRAME_GOAWAY: u64 = 0x7;
pub const FRAME_MAX_PUSH_ID: u64 = 0xd;
#[allow(dead_code)]
pub const SETTINGS_QPACK_MAX_TABLE_CAPACITY: u64 = 0x1;
#[allow(dead_code)]
pub const SETTINGS_MAX_FIELD_SECTION_SIZE: u64 = 0x6;
#[allow(dead_code)]
pub const SETTINGS_QPACK_BLOCKED_STREAMS: u64 = 0x7;
#[allow(dead_code)]
pub const SETTINGS_ENABLE_CONNECT_PROTOCOL: u64 = 0x8;
#[allow(dead_code)]
pub const SETTINGS_H3_DATAGRAM: u64 = 0x33;
pub type Setting = (u64, u64);
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct Settings {
entries: Vec<Setting>,
}
impl Settings {
#[inline]
pub fn new() -> Self {
Self::default()
}
#[inline]
pub fn insert(&mut self, id: u64, value: u64) {
self.entries.push((id, value));
}
#[inline]
pub fn get(&self, id: u64) -> Option<u64> {
self.entries
.iter()
.find_map(|(i, v)| (*i == id).then_some(*v))
}
#[inline]
pub fn iter(&self) -> impl Iterator<Item = Setting> + '_ {
self.entries.iter().copied()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Frame {
Data(Bytes),
Headers(Bytes),
Settings(Settings),
CancelPush(u64),
PushPromise { push_id: u64, field_section: Bytes },
Goaway(u64),
MaxPushId(u64),
}
impl Frame {
#[inline]
pub fn is_known(&self) -> bool {
matches!(
self,
Frame::Data(_)
| Frame::Headers(_)
| Frame::Settings(_)
| Frame::CancelPush(_)
| Frame::PushPromise { .. }
| Frame::Goaway(_)
| Frame::MaxPushId(_)
)
}
#[inline]
pub fn encode(&self, dst: &mut BytesMut) {
match self {
Frame::Data(payload) => {
write_varint(FRAME_DATA, dst);
write_varint(payload.len() as u64, dst);
dst.extend_from_slice(payload);
}
Frame::Headers(payload) => {
write_varint(FRAME_HEADERS, dst);
write_varint(payload.len() as u64, dst);
dst.extend_from_slice(payload);
}
Frame::Settings(settings) => {
write_varint(FRAME_SETTINGS, dst);
let len: usize = settings
.entries
.iter()
.map(|(id, value)| varint_size(*id) + varint_size(*value))
.sum();
write_varint(len as u64, dst);
for (id, value) in &settings.entries {
write_varint(*id, dst);
write_varint(*value, dst);
}
}
Frame::CancelPush(push_id) => {
write_varint(FRAME_CANCEL_PUSH, dst);
write_varint(varint_size(*push_id) as u64, dst);
write_varint(*push_id, dst);
}
Frame::PushPromise {
push_id,
field_section,
} => {
write_varint(FRAME_PUSH_PROMISE, dst);
write_varint((varint_size(*push_id) + field_section.len()) as u64, dst);
write_varint(*push_id, dst);
dst.extend_from_slice(field_section);
}
Frame::Goaway(stream_id) => {
write_varint(FRAME_GOAWAY, dst);
write_varint(varint_size(*stream_id) as u64, dst);
write_varint(*stream_id, dst);
}
Frame::MaxPushId(push_id) => {
write_varint(FRAME_MAX_PUSH_ID, dst);
write_varint(varint_size(*push_id) as u64, dst);
write_varint(*push_id, dst);
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FrameError {
Frame,
Unexpected(u64),
Settings,
}
impl FrameError {
pub const fn h3_code(self) -> u64 {
use crate::h3::H3Error;
match self {
FrameError::Frame => H3Error::FrameError.code(),
FrameError::Unexpected(_) => H3Error::FrameUnexpected.code(),
FrameError::Settings => H3Error::Settings.code(),
}
}
}
#[derive(Debug, Default)]
pub struct FrameDecoder {
buf: Bytes,
}
impl FrameDecoder {
#[inline]
pub fn new() -> Self {
Self::default()
}
#[inline]
pub fn extend(&mut self, data: Bytes) {
if data.is_empty() {
return;
}
if self.buf.is_empty() {
self.buf = data;
return;
}
let mut joined = BytesMut::with_capacity(self.buf.len() + data.len());
joined.extend_from_slice(&self.buf);
joined.extend_from_slice(&data);
self.buf = joined.freeze();
}
#[inline]
pub fn buffered(&self) -> usize {
self.buf.len()
}
#[inline]
pub fn next_frame(&mut self) -> Result<Option<Frame>, FrameError> {
loop {
let Some((ty, type_len)) = parse_varint(&self.buf)? else {
return Ok(None);
};
if matches!(ty, 0x02 | 0x06 | 0x08 | 0x09) {
return Err(FrameError::Unexpected(ty));
}
let Some((len, len_len)) = parse_varint(&self.buf[type_len..])? else {
return Ok(None);
};
let header_len = type_len + len_len;
let Some(total) = header_len.checked_add(len as usize) else {
return Err(FrameError::Frame);
};
if total > self.buf.len() {
return Ok(None);
}
if !is_known_frame_type(ty) {
self.buf.advance(total);
continue;
}
let mut chunk = self.buf.split_to(total);
let payload = chunk.split_off(header_len);
let frame = match ty {
FRAME_DATA => Frame::Data(payload),
FRAME_HEADERS => Frame::Headers(payload),
FRAME_CANCEL_PUSH => Frame::CancelPush(take_varint(&payload)?),
FRAME_SETTINGS => Frame::Settings(parse_settings(&payload)?),
FRAME_PUSH_PROMISE => {
let Some((push_id, id_len)) = parse_varint(&payload)? else {
return Err(FrameError::Frame);
};
Frame::PushPromise {
push_id,
field_section: payload.slice(id_len..),
}
}
FRAME_GOAWAY => Frame::Goaway(take_varint(&payload)?),
FRAME_MAX_PUSH_ID => Frame::MaxPushId(take_varint(&payload)?),
_ => unreachable!("unknown types are skipped above"),
};
return Ok(Some(frame));
}
}
}
impl FrameDecoder {
#[inline]
pub fn peek_frame_type(&self) -> Option<u64> {
match parse_varint(&self.buf) {
Ok(Some((ty, _))) => Some(ty),
_ => None,
}
}
}
#[inline]
fn is_known_frame_type(ty: u64) -> bool {
matches!(
ty,
FRAME_DATA
| FRAME_HEADERS
| FRAME_CANCEL_PUSH
| FRAME_SETTINGS
| FRAME_PUSH_PROMISE
| FRAME_GOAWAY
| FRAME_MAX_PUSH_ID
)
}
#[inline]
fn is_grease(v: u64) -> bool {
v >= 0x21 && (v - 0x21).is_multiple_of(0x1f)
}
#[inline]
fn is_reserved_setting(id: u64) -> bool {
(0x02..=0x05).contains(&id)
}
#[inline]
fn parse_settings(payload: &[u8]) -> Result<Settings, FrameError> {
let mut settings = Settings::new();
let mut rest = payload;
while !rest.is_empty() {
let Some((id, id_len)) = parse_varint(rest)? else {
return Err(FrameError::Frame);
};
rest = &rest[id_len..];
let Some((value, value_len)) = parse_varint(rest)? else {
return Err(FrameError::Frame);
};
rest = &rest[value_len..];
if is_reserved_setting(id) {
return Err(FrameError::Settings);
}
if settings.get(id).is_some() {
return Err(FrameError::Settings);
}
if !is_grease(id) {
settings.insert(id, value);
}
}
Ok(settings)
}
#[inline]
fn take_varint(buf: &[u8]) -> Result<u64, FrameError> {
let Some((value, n)) = parse_varint(buf)? else {
return Err(FrameError::Frame);
};
if n != buf.len() {
return Err(FrameError::Frame);
}
Ok(value)
}
const MIN_VARINT: [u64; 4] = [0, 1 << 6, 1 << 14, 1 << 30];
#[inline]
pub(crate) fn parse_varint(buf: &[u8]) -> Result<Option<(u64, usize)>, FrameError> {
let Some(&first) = buf.first() else {
return Ok(None);
};
let len = 1usize << (first >> 6);
if buf.len() < len {
return Ok(None);
}
let mut value = u64::from(first & 0x3f);
for &byte in &buf[1..len] {
value = (value << 8) | u64::from(byte);
}
if len > 1 && value < MIN_VARINT[usize::from(first >> 6)] {
return Err(FrameError::Frame);
}
Ok(Some((value, len)))
}
#[inline]
pub fn varint_size(value: u64) -> usize {
if value < (1 << 6) {
1
} else if value < (1 << 14) {
2
} else if value < (1 << 30) {
4
} else {
8
}
}
#[inline]
pub fn write_varint(value: u64, dst: &mut BytesMut) {
debug_assert!(value <= MAX_VARINT, "varint out of range: {value:#x}");
if value < (1 << 6) {
dst.put_u8(value as u8);
} else if value < (1 << 14) {
dst.put_u16((0b01 << 14) | value as u16);
} else if value < (1 << 30) {
dst.put_u32((0b10 << 30) | value as u32);
} else {
dst.put_u64((0b11 << 62) | value);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[inline]
fn decode_all(decoder: &mut FrameDecoder) -> Result<Vec<Frame>, FrameError> {
let mut frames = Vec::new();
while let Some(frame) = decoder.next_frame()? {
frames.push(frame);
}
Ok(frames)
}
#[inline]
fn encode_frames(frames: &[Frame]) -> Bytes {
let mut buf = BytesMut::new();
for frame in frames {
frame.encode(&mut buf);
}
buf.freeze()
}
#[test]
fn round_trip_all_frame_types() {
let mut settings = Settings::new();
settings.insert(SETTINGS_QPACK_MAX_TABLE_CAPACITY, 4096);
settings.insert(SETTINGS_MAX_FIELD_SECTION_SIZE, 100);
settings.insert(SETTINGS_QPACK_BLOCKED_STREAMS, 2);
settings.insert(0x21, 7);
let frames = [
Frame::Data(Bytes::from_static(b"hello world")),
Frame::Headers(Bytes::from_static(b"\x3f\xbd\x01")),
Frame::Settings(settings),
Frame::CancelPush(7),
Frame::PushPromise {
push_id: 1,
field_section: Bytes::from_static(b"\x05\x00\x80"),
},
Frame::Goaway(2),
Frame::MaxPushId(0),
Frame::Data(Bytes::new()),
Frame::Headers(Bytes::new()),
];
let mut decoder = FrameDecoder::new();
decoder.extend(encode_frames(&frames));
let got = decode_all(&mut decoder).expect("all frames parse");
let mut expected = frames.to_vec();
expected[2] = Frame::Settings(settings_without_grease());
assert_eq!(got, expected);
}
#[inline]
fn settings_without_grease() -> Settings {
let mut s = Settings::new();
s.insert(SETTINGS_QPACK_MAX_TABLE_CAPACITY, 4096);
s.insert(SETTINGS_MAX_FIELD_SECTION_SIZE, 100);
s.insert(SETTINGS_QPACK_BLOCKED_STREAMS, 2);
s
}
#[test]
fn incremental_byte_at_a_time() {
let wire = encode_frames(&[
Frame::Headers(Bytes::from_static(b"abc")),
Frame::Data(Bytes::from_static(b"xy")),
]);
let mut decoder = FrameDecoder::new();
let mut got = Vec::new();
for (i, &byte) in wire.iter().enumerate() {
decoder.extend(Bytes::copy_from_slice(&[byte]));
while let Some(frame) = decoder.next_frame().unwrap() {
got.push(frame);
}
if i < 4 {
assert!(got.is_empty(), "frame appeared early at byte {i}");
}
if i == 4 {
assert_eq!(got, vec![Frame::Headers(Bytes::from_static(b"abc"))]);
}
}
assert_eq!(
got,
vec![
Frame::Headers(Bytes::from_static(b"abc")),
Frame::Data(Bytes::from_static(b"xy")),
]
);
assert_eq!(decoder.buffered(), 0);
}
#[test]
fn truncated_prefixes_are_incomplete() {
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x00]));
assert_eq!(decoder.next_frame().unwrap(), None);
decoder.extend(Bytes::from_static(&[0x05]));
assert_eq!(decoder.next_frame().unwrap(), None);
decoder.extend(Bytes::from_static(b"he"));
assert_eq!(decoder.next_frame().unwrap(), None);
decoder.extend(Bytes::from_static(b"llo"));
assert_eq!(
decode_all(&mut decoder).unwrap(),
vec![Frame::Data(Bytes::from_static(b"hello"))]
);
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[
0x01, 0xc0, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff,
]));
assert_eq!(decoder.next_frame().unwrap(), None);
assert_eq!(decoder.buffered(), 10);
}
#[test]
fn forbidden_http2_frames() {
for ty in [0x02u8, 0x06, 0x08, 0x09] {
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::copy_from_slice(&[ty, 0x01, 0x00]));
let err = decoder.next_frame().unwrap_err();
assert_eq!(err, FrameError::Unexpected(u64::from(ty)));
assert_eq!(err.h3_code(), 0x0105);
}
}
#[test]
fn unknown_and_grease_frames_are_skipped() {
let mut wire = BytesMut::new();
Frame::Headers(Bytes::from_static(b"first")).encode(&mut wire);
write_varint(0x42, &mut wire);
write_varint(3, &mut wire);
wire.extend_from_slice(b"xyz");
write_varint(0x21, &mut wire);
write_varint(2, &mut wire);
wire.extend_from_slice(&[0xde, 0xad]);
Frame::Data(Bytes::from_static(b"last")).encode(&mut wire);
let mut decoder = FrameDecoder::new();
decoder.extend(wire.freeze());
let frames = decode_all(&mut decoder).unwrap();
assert_eq!(
frames,
vec![
Frame::Headers(Bytes::from_static(b"first")),
Frame::Data(Bytes::from_static(b"last")),
]
);
assert_eq!(decoder.buffered(), 0);
}
#[test]
fn settings_payload_validation() {
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x04, 0x00]));
assert_eq!(
decode_all(&mut decoder).unwrap(),
vec![Frame::Settings(Settings::new())]
);
let mut wire = BytesMut::new();
Frame::Settings({
let mut s = Settings::new();
s.insert(0x06, 1);
s.insert(0x06, 2);
s
})
.encode(&mut wire);
let mut decoder = FrameDecoder::new();
decoder.extend(wire.freeze());
let err = decoder.next_frame().unwrap_err();
assert_eq!(err, FrameError::Settings);
assert_eq!(err.h3_code(), 0x0109);
for id in 0x02..=0x05 {
let mut wire = BytesMut::new();
Frame::Settings({
let mut s = Settings::new();
s.insert(id, 0);
s
})
.encode(&mut wire);
let mut decoder = FrameDecoder::new();
decoder.extend(wire.freeze());
assert_eq!(decoder.next_frame().unwrap_err(), FrameError::Settings);
}
let mut s = Settings::new();
s.insert(0x0100, 5);
s.insert(SETTINGS_ENABLE_CONNECT_PROTOCOL, 1);
let mut wire = BytesMut::new();
Frame::Settings(s.clone()).encode(&mut wire);
let mut decoder = FrameDecoder::new();
decoder.extend(wire.freeze());
match decode_all(&mut decoder).unwrap()[0].clone() {
Frame::Settings(got) => {
assert_eq!(got.get(0x0100), Some(5));
assert_eq!(got.get(SETTINGS_ENABLE_CONNECT_PROTOCOL), Some(1));
}
other => panic!("expected Settings, got {other:?}"),
}
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x04, 0x01, 0x06]));
assert_eq!(decoder.next_frame().unwrap_err(), FrameError::Frame);
}
#[test]
fn fixed_value_frames_reject_bad_payloads() {
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x03, 0x00]));
assert_eq!(decoder.next_frame().unwrap_err(), FrameError::Frame);
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x03, 0x02, 0x01, 0x00]));
assert_eq!(decoder.next_frame().unwrap_err(), FrameError::Frame);
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x03, 0x02, 0x40, 0x05]));
assert_eq!(decoder.next_frame().unwrap_err(), FrameError::Frame);
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x07, 0x01, 0x05]));
assert_eq!(decode_all(&mut decoder).unwrap(), vec![Frame::Goaway(5)]);
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x40, 0x00, 0x00]));
assert_eq!(decoder.next_frame().unwrap_err(), FrameError::Frame);
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x00, 0x40, 0x00]));
assert_eq!(decoder.next_frame().unwrap_err(), FrameError::Frame);
}
#[test]
fn push_promise_shapes() {
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x05, 0x04, 0x01, b'a', b'b', b'c']));
assert_eq!(
decode_all(&mut decoder).unwrap(),
vec![Frame::PushPromise {
push_id: 1,
field_section: Bytes::from_static(b"abc"),
}]
);
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x05, 0x01, 0x01]));
assert_eq!(
decode_all(&mut decoder).unwrap(),
vec![Frame::PushPromise {
push_id: 1,
field_section: Bytes::new(),
}]
);
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x05, 0x00]));
assert_eq!(decoder.next_frame().unwrap_err(), FrameError::Frame);
}
#[test]
fn varint_edge_encodings() {
for value in [
(1 << 6) - 1,
1 << 6,
(1 << 14) - 1,
1 << 14,
(1 << 30) - 1,
1 << 30,
MAX_VARINT,
] {
let mut wire = BytesMut::new();
write_varint(value, &mut wire);
assert_eq!(wire.len(), varint_size(value));
let (got, n) = parse_varint(&wire).unwrap().unwrap();
assert_eq!(got, value);
assert_eq!(n, wire.len());
}
assert_eq!(parse_varint(&[0x40, 0x00]), Err(FrameError::Frame)); assert_eq!(
parse_varint(&[0x80, 0x00, 0x00, 0x40]),
Err(FrameError::Frame)
); assert_eq!(parse_varint(&[0x40, 0x40]), Ok(Some((64, 2))));
assert_eq!(parse_varint(&[0x40]), Ok(None));
assert_eq!(parse_varint(&[]), Ok(None));
}
#[test]
fn clean_eof_with_truncated_frame_is_detectable() {
let mut decoder = FrameDecoder::new();
decoder.extend(Bytes::from_static(&[0x01, 0x05, b'a', b'b']));
assert_eq!(decoder.next_frame().unwrap(), None);
assert_eq!(decoder.buffered(), 4);
}
}