use crate::{
Buf,
BufError::{self},
BufMut, BufResult, Codec, Cursor,
ietf::quic::{HeaderForm, Version},
ietf::quicv1::{
ConnectionId, FixedBit, KeyPhase, Length, PacketNumber, PacketNumberLength, RetryToken,
SpinBit, VariableLengthInteger,
},
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
#[repr(u8)]
pub enum LongHeaderPacketType {
Initial = 0b01,
ZeroRtt = 0b10,
Handshake = 0b11,
Retry = 0b00,
}
impl Codec for LongHeaderPacketType {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, _: ()) -> BufResult<()> {
let current_byte = writer.peek_u8().unwrap_or(0x00);
let bit_mask = (*self as u8) << 4;
let updated_byte = (current_byte & !0x30) | bit_mask;
writer.poke_u8(updated_byte)
}
fn decode<R: Buf>(reader: &mut Cursor<R>, _: ()) -> BufResult<Self> {
let value = (reader.peek_u8()? & 0x30) >> 4;
match value {
0 => Ok(Self::Initial),
1 => Ok(Self::ZeroRtt),
2 => Ok(Self::Handshake),
_ => Ok(Self::Retry),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RetryIntegrityTag(pub [u8; 16]);
impl Codec for RetryIntegrityTag {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, _: ()) -> BufResult<()> {
writer.write_array(&self.0)
}
fn decode<R: Buf>(reader: &mut Cursor<R>, _: ()) -> BufResult<Self> {
Ok(Self(reader.read_array::<16>()?))
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct InitialPacket {
pub version: Version,
pub destination_connection_id: ConnectionId,
pub source_connection_id: ConnectionId,
pub token: RetryToken,
pub packet_number: PacketNumber,
pub packet_payload: Vec<u8>,
}
impl InitialPacket {
pub const HEADER_FORM: HeaderForm = HeaderForm::LongHeader;
pub const FIXED_BIT: FixedBit = FixedBit::One;
pub const LONG_PACKET_TYPE: LongHeaderPacketType = LongHeaderPacketType::Initial;
}
impl Codec for InitialPacket {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, _: ()) -> BufResult<()> {
Self::HEADER_FORM.encode(writer, ())?;
Self::FIXED_BIT.encode(writer, ())?;
Self::LONG_PACKET_TYPE.encode(writer, ())?;
let packet_number_length = self.packet_number.length()?;
packet_number_length.encode(writer, ())?;
self.version.encode(writer, ())?;
self.destination_connection_id.encode(writer, ())?;
self.source_connection_id.encode(writer, ())?;
VariableLengthInteger::new(self.token.0.len() as u64)?.encode(writer, ())?;
self.token.0.encode(writer, ())?;
Length::calculate(self.packet_payload.len(), packet_number_length)?.encode(writer, ())?;
self.packet_number.encode(writer, packet_number_length)?;
self.packet_payload.encode(writer, ())
}
fn decode<R: Buf>(reader: &mut Cursor<R>, _: ()) -> BufResult<Self> {
if HeaderForm::decode(reader, ())? != Self::HEADER_FORM {
return Err(BufError::UnexpectedValue);
}
if FixedBit::decode(reader, ())? != Self::FIXED_BIT {
return Err(BufError::UnexpectedValue);
}
if LongHeaderPacketType::decode(reader, ())? != Self::LONG_PACKET_TYPE {
return Err(BufError::UnexpectedValue);
}
let packet_number_length = PacketNumberLength::decode(reader, ())?;
let version = Version::decode(reader, ())?;
let destination_connection_id = ConnectionId::decode(reader, ())?;
let source_connection_id = ConnectionId::decode(reader, ())?;
let token_length = VariableLengthInteger::decode(reader, ())?;
let token_bytes = &mut [0u8; 16];
reader.read_into(&mut token_bytes[..token_length.0 as usize])?;
let token = RetryToken((&token_bytes[..token_length.0 as usize]).to_vec());
let _length = VariableLengthInteger::decode(reader, ())?;
let packet_number = PacketNumber::decode(reader, packet_number_length)?;
let packet_payload = Vec::decode(reader, ())?;
Ok(Self {
version,
destination_connection_id,
source_connection_id,
token,
packet_number,
packet_payload,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ZeroRttPacket {
pub version: Version,
pub destination_connection_id: ConnectionId,
pub source_connection_id: ConnectionId,
pub packet_number: PacketNumber,
pub packet_payload: Vec<u8>,
}
impl ZeroRttPacket {
pub const HEADER_FORM: HeaderForm = HeaderForm::LongHeader;
pub const FIXED_BIT: FixedBit = FixedBit::One;
pub const LONG_PACKET_TYPE: LongHeaderPacketType = LongHeaderPacketType::ZeroRtt;
}
impl Codec for ZeroRttPacket {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, _: ()) -> BufResult<()> {
Self::HEADER_FORM.encode(writer, ())?;
Self::FIXED_BIT.encode(writer, ())?;
Self::LONG_PACKET_TYPE.encode(writer, ())?;
let packet_number_length = self.packet_number.length()?;
packet_number_length.encode(writer, ())?;
self.version.encode(writer, ())?;
self.destination_connection_id.encode(writer, ())?;
self.source_connection_id.encode(writer, ())?;
Length::calculate(self.packet_payload.len(), packet_number_length)?.encode(writer, ())?;
self.packet_number.encode(writer, packet_number_length)?;
self.packet_payload.encode(writer, ())
}
fn decode<R: Buf>(reader: &mut Cursor<R>, _: ()) -> BufResult<Self> {
if HeaderForm::decode(reader, ())? != Self::HEADER_FORM {
return Err(BufError::UnexpectedValue);
}
if FixedBit::decode(reader, ())? != Self::FIXED_BIT {
return Err(BufError::UnexpectedValue);
}
if LongHeaderPacketType::decode(reader, ())? != Self::LONG_PACKET_TYPE {
return Err(BufError::UnexpectedValue);
}
let packet_number_length = PacketNumberLength::decode(reader, ())?;
let version = Version::decode(reader, ())?;
let destination_connection_id = ConnectionId::decode(reader, ())?;
let source_connection_id = ConnectionId::decode(reader, ())?;
let _length = VariableLengthInteger::decode(reader, ())?;
let packet_number = PacketNumber::decode(reader, packet_number_length)?;
let packet_payload = Vec::decode(reader, ())?;
Ok(Self {
version,
destination_connection_id,
source_connection_id,
packet_number,
packet_payload,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct HandshakePacket {
pub version: Version,
pub destination_connection_id: ConnectionId,
pub source_connection_id: ConnectionId,
pub packet_number: PacketNumber,
pub packet_payload: Vec<u8>,
}
impl HandshakePacket {
pub const HEADER_FORM: HeaderForm = HeaderForm::LongHeader;
pub const FIXED_BIT: FixedBit = FixedBit::One;
pub const LONG_PACKET_TYPE: LongHeaderPacketType = LongHeaderPacketType::Handshake;
}
impl Codec for HandshakePacket {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, _: ()) -> BufResult<()> {
Self::HEADER_FORM.encode(writer, ())?;
Self::FIXED_BIT.encode(writer, ())?;
Self::LONG_PACKET_TYPE.encode(writer, ())?;
let packet_number_length = self.packet_number.length()?;
packet_number_length.encode(writer, ())?;
self.version.encode(writer, ())?;
self.destination_connection_id.encode(writer, ())?;
self.source_connection_id.encode(writer, ())?;
Length::calculate(self.packet_payload.len(), packet_number_length)?.encode(writer, ())?;
self.packet_number.encode(writer, packet_number_length)?;
self.packet_payload.encode(writer, ())
}
fn decode<R: Buf>(reader: &mut Cursor<R>, _: ()) -> BufResult<Self> {
if HeaderForm::decode(reader, ())? != Self::HEADER_FORM {
return Err(BufError::UnexpectedValue);
}
if FixedBit::decode(reader, ())? != Self::FIXED_BIT {
return Err(BufError::UnexpectedValue);
}
if LongHeaderPacketType::decode(reader, ())? != Self::LONG_PACKET_TYPE {
return Err(BufError::UnexpectedValue);
}
let packet_number_length = PacketNumberLength::decode(reader, ())?;
let version = Version::decode(reader, ())?;
let destination_connection_id = ConnectionId::decode(reader, ())?;
let source_connection_id = ConnectionId::decode(reader, ())?;
let _length = VariableLengthInteger::decode(reader, ())?;
let packet_number = PacketNumber::decode(reader, packet_number_length)?;
let packet_payload = Vec::decode(reader, ())?;
Ok(Self {
version,
destination_connection_id,
source_connection_id,
packet_number,
packet_payload,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RetryPacket {
pub version: Version,
pub destination_connection_id: ConnectionId,
pub source_connection_id: ConnectionId,
pub retry_token: RetryToken,
pub retry_integrity_tag: RetryIntegrityTag,
}
impl RetryPacket {
pub const HEADER_FORM: HeaderForm = HeaderForm::LongHeader;
pub const FIXED_BIT: FixedBit = FixedBit::One;
pub const LONG_PACKET_TYPE: LongHeaderPacketType = LongHeaderPacketType::Retry;
}
impl Codec for RetryPacket {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, _: ()) -> BufResult<()> {
Self::HEADER_FORM.encode(writer, ())?;
Self::FIXED_BIT.encode(writer, ())?;
Self::LONG_PACKET_TYPE.encode(writer, ())?;
writer.advance(1)?;
self.version.encode(writer, ())?;
self.destination_connection_id.encode(writer, ())?;
self.source_connection_id.encode(writer, ())?;
self.retry_token.encode(writer, ())?;
self.retry_integrity_tag.encode(writer, ())
}
fn decode<R: Buf>(reader: &mut Cursor<R>, _: ()) -> BufResult<Self> {
if HeaderForm::decode(reader, ())? != Self::HEADER_FORM {
return Err(BufError::UnexpectedValue);
}
if FixedBit::decode(reader, ())? != Self::FIXED_BIT {
return Err(BufError::UnexpectedValue);
}
if LongHeaderPacketType::decode(reader, ())? != Self::LONG_PACKET_TYPE {
return Err(BufError::UnexpectedValue);
}
reader.advance(1)?;
let version = Version::decode(reader, ())?;
let destination_connection_id = ConnectionId::decode(reader, ())?;
let source_connection_id = ConnectionId::decode(reader, ())?;
let mut left = Vec::decode(reader, ())?;
let length = left.len();
let (retry_token, _retry_integrity_tag) = left.split_at_mut_checked(length - 16).unwrap();
let retry_token = RetryToken(retry_token.to_vec());
let retry_integrity_tag = match _retry_integrity_tag.try_into() {
Ok(x) => x,
Err(_) => return Err(BufError::UnexpectedValue),
};
let retry_integrity_tag = RetryIntegrityTag(retry_integrity_tag);
Ok(Self {
version,
destination_connection_id,
source_connection_id,
retry_token,
retry_integrity_tag,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct OneRttPacket {
pub spin_bit: SpinBit,
pub key_phase: KeyPhase,
pub destination_connection_id: ConnectionId,
pub packet_number: PacketNumber,
pub packet_payload: Vec<u8>,
}
impl OneRttPacket {
pub const HEADER_FORM: HeaderForm = HeaderForm::ShortHeader;
pub const FIXED_BIT: FixedBit = FixedBit::One;
}
impl Codec<u8> for OneRttPacket {
fn encode<W: BufMut>(&self, writer: &mut Cursor<W>, dcid_length: u8) -> BufResult<()> {
Self::HEADER_FORM.encode(writer, ())?;
Self::FIXED_BIT.encode(writer, ())?;
self.spin_bit.encode(writer, ())?;
self.key_phase.encode(writer, ())?;
let packet_number_length = self.packet_number.length()?;
packet_number_length.encode(writer, ())?;
self.destination_connection_id.encode(writer, dcid_length)?;
self.packet_number.encode(writer, packet_number_length)?;
self.packet_payload.encode(writer, ())
}
fn decode<R: Buf>(reader: &mut Cursor<R>, dcid_length: u8) -> BufResult<Self> {
if HeaderForm::decode(reader, ())? != Self::HEADER_FORM {
return Err(BufError::UnexpectedValue);
}
if FixedBit::decode(reader, ())? != Self::FIXED_BIT {
return Err(BufError::UnexpectedValue);
}
let spin_bit = SpinBit::decode(reader, ())?;
let key_phase = KeyPhase::decode(reader, ())?;
let packet_number_length = PacketNumberLength::decode(reader, ())?;
let destination_connection_id = ConnectionId::decode(reader, dcid_length)?;
let packet_number = PacketNumber::decode(reader, packet_number_length)?;
let packet_payload = Vec::decode(reader, ())?;
Ok(Self {
spin_bit,
key_phase,
destination_connection_id,
packet_number,
packet_payload,
})
}
}