use tracing::{debug, trace};
use crate::Instant;
use crate::connection::spaces::PacketSpace;
use crate::crypto::{HeaderKey, KeyPair, PacketKey};
use crate::packet::{Packet, PartialDecode, SpaceId};
use crate::token::ResetToken;
use crate::{RESET_TOKEN_SIZE, TransportError};
pub(super) fn unprotect_header(
partial_decode: PartialDecode,
spaces: &[PacketSpace; 3],
zero_rtt_crypto: Option<&ZeroRttCrypto>,
stateless_reset_token: Option<ResetToken>,
) -> Option<UnprotectHeaderResult> {
let header_crypto = if partial_decode.is_0rtt() {
if let Some(crypto) = zero_rtt_crypto {
Some(&*crypto.header)
} else {
debug!("dropping unexpected 0-RTT packet");
return None;
}
} else if let Some(space) = partial_decode.space() {
if let Some(ref crypto) = spaces[space].crypto {
Some(&*crypto.header.remote)
} else {
debug!(
"discarding unexpected {:?} packet ({} bytes)",
space,
partial_decode.len(),
);
return None;
}
} else {
None
};
let packet = partial_decode.data();
let stateless_reset = packet.len() >= RESET_TOKEN_SIZE + 5
&& stateless_reset_token.as_deref() == Some(&packet[packet.len() - RESET_TOKEN_SIZE..]);
match partial_decode.finish(header_crypto) {
Ok(packet) => Some(UnprotectHeaderResult {
packet: Some(packet),
stateless_reset,
}),
Err(_) if stateless_reset => Some(UnprotectHeaderResult {
packet: None,
stateless_reset: true,
}),
Err(e) => {
trace!("unable to complete packet decoding: {}", e);
None
}
}
}
pub(super) struct UnprotectHeaderResult {
pub(super) packet: Option<Packet>,
pub(super) stateless_reset: bool,
}
pub(super) fn decrypt_packet_body(
packet: &mut Packet,
spaces: &[PacketSpace; 3],
zero_rtt_crypto: Option<&ZeroRttCrypto>,
conn_key_phase: bool,
prev_crypto: Option<&PrevCrypto>,
next_crypto: Option<&KeyPair<Box<dyn PacketKey>>>,
) -> Result<Option<DecryptPacketResult>, Option<TransportError>> {
if !packet.header.is_protected() {
return Ok(None);
}
let space = packet.header.space();
let rx_packet = spaces[space].rx_packet;
let number = packet.header.number().ok_or(None)?.expand(rx_packet + 1);
let packet_key_phase = packet.header.key_phase();
let mut crypto_update = false;
let crypto = if packet.header.is_0rtt() {
&zero_rtt_crypto.unwrap().packet
} else if packet_key_phase == conn_key_phase || space != SpaceId::Data {
&spaces[space].crypto.as_ref().unwrap().packet.remote
} else if let Some(prev) = prev_crypto.filter(|&crypto| {
crypto.end_packet.is_none_or(|(pn, _)| number < pn)
}) {
&prev.crypto.remote
} else {
crypto_update = true;
&next_crypto.unwrap().remote
};
crypto
.decrypt(number, &packet.header_data, &mut packet.payload)
.map_err(|_| {
trace!("decryption failed with packet number {}", number);
None
})?;
if !packet.reserved_bits_valid() {
return Err(Some(TransportError::PROTOCOL_VIOLATION(
"reserved bits set",
)));
}
let mut outgoing_key_update_acked = false;
if let Some(prev) = prev_crypto {
if prev.end_packet.is_none() && packet_key_phase == conn_key_phase {
outgoing_key_update_acked = true;
}
}
if crypto_update {
if number <= rx_packet || prev_crypto.is_some_and(|x| x.update_unacked) {
return Err(Some(TransportError::KEY_UPDATE_ERROR("")));
}
}
Ok(Some(DecryptPacketResult {
number,
outgoing_key_update_acked,
incoming_key_update: crypto_update,
}))
}
pub(super) struct DecryptPacketResult {
pub(super) number: u64,
pub(super) outgoing_key_update_acked: bool,
pub(super) incoming_key_update: bool,
}
pub(super) struct PrevCrypto {
pub(super) crypto: KeyPair<Box<dyn PacketKey>>,
pub(super) end_packet: Option<(u64, Instant)>,
pub(super) update_unacked: bool,
}
pub(super) struct ZeroRttCrypto {
pub(super) header: Box<dyn HeaderKey>,
pub(super) packet: Box<dyn PacketKey>,
}