use crate::courierust_error::{Error, Result};
use crate::courierust_tls::crypto::hash::{BoxDigest, Sha256, Sha384};
use crate::courierust_tls::crypto::hmac::{expand_label, extract};
use alloc::vec::Vec;
pub const TLS_AES_128_GCM_SHA256: u16 = 0x1301;
pub const TLS_AES_256_GCM_SHA384: u16 = 0x1302;
pub const TLS_CHACHA20_POLY1305_SHA256: u16 = 0x1303;
pub const INITIAL_SALT_V1: [u8; 20] = [
0x38, 0x76, 0x2c, 0xf7, 0xf5, 0x59, 0x34, 0xb3, 0x4d, 0x17, 0x9a, 0xe6, 0xa4, 0xc8, 0x0c, 0xad,
0xcc, 0xbb, 0x7f, 0x0a,
];
const AEAD_TAG_LEN: usize = 16;
const RETRY_INTEGRITY_KEY: [u8; 16] = [
0xbe, 0x0c, 0x69, 0x0b, 0x9f, 0x66, 0x57, 0x5a, 0x1d, 0x76, 0x6b, 0x54, 0xe3, 0x68, 0xc8, 0x4e,
];
const RETRY_INTEGRITY_NONCE: [u8; 12] = [
0x46, 0x15, 0x99, 0xd3, 0x5d, 0x63, 0x2b, 0xf2, 0x23, 0x98, 0x25, 0xbb,
];
fn digest_for_suite(suite: u16) -> BoxDigest {
match suite {
TLS_AES_256_GCM_SHA384 => Box::new(Sha384::new()),
_ => Box::new(Sha256::new()),
}
}
fn hash_len(suite: u16) -> usize {
if suite == TLS_AES_256_GCM_SHA384 {
48
} else {
32
}
}
fn key_len(suite: u16) -> Option<usize> {
match suite {
TLS_AES_128_GCM_SHA256 => Some(16),
TLS_AES_256_GCM_SHA384 | TLS_CHACHA20_POLY1305_SHA256 => Some(32),
_ => None,
}
}
fn expand_quic(suite: u16, secret: &[u8], label: &[u8], len: usize) -> Vec<u8> {
expand_label(digest_for_suite(suite).as_mut(), secret, label, &[], len)
}
fn nonce(iv: &[u8; 12], packet_number: u64) -> [u8; 12] {
let mut out = *iv;
for (slot, byte) in out[4..].iter_mut().zip(packet_number.to_be_bytes()) {
*slot ^= byte;
}
out
}
#[derive(Clone)]
pub struct PacketKey {
suite: u16,
key: [u8; 32],
iv: [u8; 12],
hp: [u8; 32],
key_len: usize,
secret: [u8; 48],
secret_len: usize,
}
impl core::fmt::Debug for PacketKey {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("PacketKey")
.field("suite", &format_args!("0x{:04x}", self.suite))
.field("key_len", &self.key_len)
.finish_non_exhaustive()
}
}
impl PacketKey {
pub fn from_secret(suite: u16, secret: &[u8]) -> Result<Self> {
let key_len =
key_len(suite).ok_or_else(|| Error::protocol("unsupported QUIC cipher suite"))?;
if secret.len() != hash_len(suite) {
return Err(Error::protocol("invalid QUIC traffic-secret length"));
}
let key = expand_quic(suite, secret, b"quic key", key_len);
let iv = expand_quic(suite, secret, b"quic iv", 12);
let hp = expand_quic(suite, secret, b"quic hp", key_len);
let mut key_arr = [0u8; 32];
let mut hp_arr = [0u8; 32];
let mut iv_arr = [0u8; 12];
let mut secret_arr = [0u8; 48];
key_arr[..key_len].copy_from_slice(&key);
hp_arr[..key_len].copy_from_slice(&hp);
iv_arr.copy_from_slice(&iv);
secret_arr[..secret.len()].copy_from_slice(secret);
Ok(Self {
suite,
key: key_arr,
iv: iv_arr,
hp: hp_arr,
key_len,
secret: secret_arr,
secret_len: secret.len(),
})
}
pub fn next_key_phase(&self) -> Result<Self> {
let secret = expand_quic(
self.suite,
&self.secret[..self.secret_len],
b"quic ku",
self.secret_len,
);
let mut next = Self::from_secret(self.suite, &secret)?;
next.hp = self.hp;
Ok(next)
}
pub fn initial(dcid: &[u8], server_direction: bool) -> Result<Self> {
if dcid.is_empty() || dcid.len() > 20 {
return Err(Error::protocol("QUIC Initial DCID must contain 1-20 bytes"));
}
let mut digest = Sha256::new();
let initial_secret = extract(&mut digest, &INITIAL_SALT_V1, dcid);
let label = if server_direction {
b"server in"
} else {
b"client in"
};
let secret = expand_quic(TLS_AES_128_GCM_SHA256, &initial_secret, label, 32);
Self::from_secret(TLS_AES_128_GCM_SHA256, &secret)
}
pub fn suite(&self) -> u16 {
self.suite
}
pub fn fingerprint(&self) -> [u8; 8] {
let mut out = [0u8; 8];
out[..4].copy_from_slice(&self.key[..4]);
out[4..].copy_from_slice(&self.iv[..4]);
out
}
pub fn seal(&self, packet_number: u64, header: &[u8], plaintext: &[u8]) -> Result<Vec<u8>> {
let iv = nonce(&self.iv, packet_number);
let sealed = match self.suite {
TLS_AES_128_GCM_SHA256 | TLS_AES_256_GCM_SHA384 => {
crate::courierust_tls::crypto::gcm::seal(
&self.key[..self.key_len],
&iv,
header,
plaintext,
)
}
TLS_CHACHA20_POLY1305_SHA256 => {
let key: &[u8; 32] = self.key[..32]
.try_into()
.map_err(|_| Error::protocol("invalid ChaCha20 key length"))?;
Some(crate::courierust_tls::crypto::chacha20poly1305::seal(
key, &iv, header, plaintext,
))
}
_ => None,
}
.ok_or_else(|| Error::protocol("QUIC AEAD seal failed"))?;
Ok(sealed)
}
pub fn open(&self, packet_number: u64, header: &[u8], sealed: &[u8]) -> Result<Vec<u8>> {
if sealed.len() < AEAD_TAG_LEN {
return Err(Error::protocol("QUIC packet is shorter than its AEAD tag"));
}
let iv = nonce(&self.iv, packet_number);
let plain = match self.suite {
TLS_AES_128_GCM_SHA256 | TLS_AES_256_GCM_SHA384 => {
crate::courierust_tls::crypto::gcm::open(
&self.key[..self.key_len],
&iv,
header,
sealed,
)
}
TLS_CHACHA20_POLY1305_SHA256 => {
let key: &[u8; 32] = self.key[..32]
.try_into()
.map_err(|_| Error::protocol("invalid ChaCha20 key length"))?;
crate::courierust_tls::crypto::chacha20poly1305::open(key, &iv, header, sealed)
}
_ => None,
}
.ok_or_else(|| Error::protocol("QUIC packet authentication failed"))?;
Ok(plain)
}
fn hp_mask(&self, sample: &[u8; 16]) -> [u8; 5] {
let mut mask = [0u8; 5];
match self.suite {
TLS_AES_128_GCM_SHA256 | TLS_AES_256_GCM_SHA384 => {
let aes = crate::courierust_tls::crypto::aes::Aes::new(&self.hp[..self.key_len])
.expect("validated AES header-protection key");
let mut block = *sample;
aes.encrypt_block(&mut block);
mask.copy_from_slice(&block[..5]);
}
TLS_CHACHA20_POLY1305_SHA256 => {
let key: &[u8; 32] = self.hp[..32].try_into().expect("validated ChaCha key");
let nonce: &[u8; 12] = sample[4..].try_into().expect("QUIC HP sample size");
let counter = u32::from_le_bytes(sample[..4].try_into().expect("QUIC HP sample"));
let chacha = crate::courierust_tls::crypto::chacha20::ChaCha20::new(key, nonce);
let mut block = [0u8; 64];
chacha.block_at(counter, &mut block);
mask.copy_from_slice(&block[..5]);
}
_ => unreachable!("PacketKey validates the suite"),
}
mask
}
pub fn protect_header(
&self,
packet: &mut [u8],
pn_offset: usize,
long_header: bool,
) -> Result<()> {
let sample_start = pn_offset
.checked_add(4)
.ok_or_else(|| Error::overflow("QUIC header-protection sample offset overflow"))?;
let sample_end = sample_start
.checked_add(16)
.ok_or_else(|| Error::overflow("QUIC header-protection sample end overflow"))?;
let sample = packet
.get(sample_start..sample_end)
.ok_or_else(|| Error::protocol("QUIC packet is too short for header protection"))?;
let sample: &[u8; 16] = sample.try_into().expect("checked QUIC sample length");
let mask = self.hp_mask(sample);
let pn_len = (packet[0] & 0x03) as usize + 1;
packet[0] ^= mask[0] & if long_header { 0x0f } else { 0x1f };
let pn_end = pn_offset
.checked_add(pn_len)
.ok_or_else(|| Error::overflow("QUIC packet number end overflow"))?;
if packet.len() < pn_end {
return Err(Error::protocol("QUIC packet number is truncated"));
}
for (byte, m) in packet[pn_offset..pn_end]
.iter_mut()
.zip(mask.iter().skip(1))
{
*byte ^= *m;
}
Ok(())
}
pub fn unprotect_header(
&self,
packet: &mut [u8],
pn_offset: usize,
long_header: bool,
) -> Result<usize> {
let sample_start = pn_offset
.checked_add(4)
.ok_or_else(|| Error::overflow("QUIC header-protection sample offset overflow"))?;
let sample_end = sample_start
.checked_add(16)
.ok_or_else(|| Error::overflow("QUIC header-protection sample end overflow"))?;
let sample = packet
.get(sample_start..sample_end)
.ok_or_else(|| Error::protocol("QUIC packet is too short for header protection"))?;
let sample: &[u8; 16] = sample.try_into().expect("checked QUIC sample length");
let mask = self.hp_mask(sample);
packet[0] ^= mask[0] & if long_header { 0x0f } else { 0x1f };
let pn_len = (packet[0] & 0x03) as usize + 1;
let pn_end = pn_offset
.checked_add(pn_len)
.ok_or_else(|| Error::overflow("QUIC packet number end overflow"))?;
if packet.len() < pn_end {
return Err(Error::protocol("QUIC packet number is truncated"));
}
for (byte, m) in packet[pn_offset..pn_end]
.iter_mut()
.zip(mask.iter().skip(1))
{
*byte ^= *m;
}
if packet[0] & 0x40 == 0 || packet[0] & if long_header { 0x0c } else { 0x18 } != 0 {
return Err(Error::protocol(
"QUIC fixed or reserved header bits are invalid",
));
}
Ok(pn_len)
}
}
pub fn initial_pair(dcid: &[u8]) -> Result<(PacketKey, PacketKey)> {
Ok((
PacketKey::initial(dcid, false)?,
PacketKey::initial(dcid, true)?,
))
}
pub fn retry_integrity_tag(original_dcid: &[u8], retry_packet: &[u8]) -> Result<[u8; 16]> {
if original_dcid.len() > 20 {
return Err(Error::protocol(
"QUIC original DCID is longer than 20 bytes",
));
}
let mut aad = Vec::with_capacity(1 + original_dcid.len() + retry_packet.len());
aad.push(original_dcid.len() as u8);
aad.extend_from_slice(original_dcid);
aad.extend_from_slice(retry_packet);
let tag = crate::courierust_tls::crypto::gcm::seal(
&RETRY_INTEGRITY_KEY,
&RETRY_INTEGRITY_NONCE,
&aad,
&[],
)
.ok_or_else(|| Error::protocol("QUIC Retry integrity calculation failed"))?;
tag.as_slice()
.try_into()
.map_err(|_| Error::protocol("QUIC Retry integrity tag has invalid length"))
}
pub fn verify_retry_integrity(
original_dcid: &[u8],
retry_packet_without_tag: &[u8],
tag: &[u8],
) -> Result<bool> {
if tag.len() != AEAD_TAG_LEN {
return Ok(false);
}
let expected = retry_integrity_tag(original_dcid, retry_packet_without_tag)?;
let mut difference = 0u8;
for (left, right) in expected.iter().zip(tag) {
difference |= left ^ right;
}
Ok(difference == 0)
}
#[cfg(test)]
mod tests {
use super::*;
fn hex(d: &[u8]) -> String {
let mut s = String::new();
for b in d {
s.push_str(&format!("{b:02x}"));
}
s
}
#[test]
fn initial_keys_match_rfc9001_appendix_a1() {
let dcid = [0x83, 0x94, 0xc8, 0xf0, 0x3e, 0x51, 0x57, 0x08];
let (client, server) = initial_pair(&dcid).unwrap();
assert_eq!(
hex(&client.fingerprint()),
hex(&[
0x1f, 0x36, 0x96, 0x13, 0xfa, 0x04, 0x4b, 0x2f, ]),
"client Initial keys diverge from RFC 9001 A.1 (client fp={}, server fp={})",
hex(&client.fingerprint()),
hex(&server.fingerprint())
);
assert_eq!(
hex(&server.fingerprint()),
hex(&[
0xcf, 0x3a, 0x53, 0x31, 0x0a, 0xc1, 0x49, 0x3c, ]),
"server Initial keys diverge from RFC 9001 A.1"
);
let mut header = vec![
0xc1, 0, 0, 0, 1, 8, 0x83, 0x94, 0xc8, 0xf0, 0x3e, 0x51, 0x57, 0x08, 0, 0, 0x40, 0x04,
0, 1,
];
let plain = b"ping";
let sealed = client.seal(0, &header, plain).unwrap();
header.extend_from_slice(&sealed);
assert_eq!(
client
.open(0, &header[..header.len() - sealed.len()], &sealed)
.unwrap(),
plain
);
}
#[test]
fn key_update_matches_rfc9001_appendix_a5() {
let secret: [u8; 32] = [
0x9a, 0xc3, 0x12, 0xa7, 0xf8, 0x77, 0x46, 0x8e, 0xbe, 0x69, 0x42, 0x27, 0x48, 0xad,
0x00, 0xa1, 0x54, 0x43, 0xf1, 0x82, 0x03, 0xa0, 0x7d, 0x60, 0x60, 0xf6, 0x88, 0xf3,
0x0f, 0x21, 0x63, 0x2b,
];
let key = PacketKey::from_secret(TLS_CHACHA20_POLY1305_SHA256, &secret).unwrap();
assert_eq!(
hex(&key.fingerprint()),
hex(&[
0xc6, 0xd9, 0x8f, 0xf3, 0xe0, 0x45, 0x9b, 0x34, ]),
"quic key/iv diverge from RFC 9001 A.5"
);
let ku_secret = expand_quic(TLS_CHACHA20_POLY1305_SHA256, &secret, b"quic ku", 32);
assert_eq!(
hex(&ku_secret),
"1223504755036d556342ee9361d253421a826c9ecdf3c7148684b36b714881f9",
"quic ku secret diverges from RFC 9001 A.5"
);
let next = key.next_key_phase().unwrap();
let expected_next_key =
expand_quic(TLS_CHACHA20_POLY1305_SHA256, &ku_secret, b"quic key", 32);
assert_eq!(hex(&next.fingerprint()[..4]), hex(&expected_next_key[..4]));
let mut header = [0u8; 13];
header[0] = 0x40 | 0x04 | 0x03;
header[1..9].copy_from_slice(&[9, 9, 9, 9, 9, 9, 9, 9]);
header[9..13].copy_from_slice(&7u32.to_be_bytes());
let sealed = next.seal(7, &header, b"ku-test").unwrap();
let mut wire = header.to_vec();
wire.extend_from_slice(&sealed);
next.protect_header(&mut wire, 9, false).unwrap();
let pn_len = next.unprotect_header(&mut wire, 9, false).unwrap();
let pn = crate::courierust_quic::packet::decode_pn(&wire[9..9 + pn_len], 7, pn_len);
assert_eq!(pn, 7);
let plain = next
.open(pn, &wire[..9 + pn_len], &wire[9 + pn_len..])
.unwrap();
assert_eq!(plain, b"ku-test");
}
#[test]
fn initial_keys_are_directional() {
let (client, server) =
initial_pair(&[0x83, 0x94, 0xc8, 0xf0, 0x3e, 0x51, 0x57, 0x08]).unwrap();
let mut header = vec![
0xc1, 0, 0, 0, 1, 8, 0x83, 0x94, 0xc8, 0xf0, 0x3e, 0x51, 0x57, 0x08, 0, 0, 0x40, 0x04,
0, 1,
];
let plain = b"ping";
let sealed = client.seal(0, &header, plain).unwrap();
header.extend_from_slice(&sealed);
assert_eq!(
client
.open(0, &header[..header.len() - sealed.len()], &sealed)
.unwrap(),
plain
);
assert!(server
.open(0, &header[..header.len() - sealed.len()], &sealed)
.is_err());
assert!(server
.open(1, &header[..header.len() - sealed.len()], &sealed)
.is_err());
}
}