From 28aa2f1f9bee0234e53653b19089645a08e570ff Mon Sep 17 00:00:00 2001 From: Benjamin Saunders Date: Sun, 23 Dec 2018 14:09:35 -0800 Subject: [PATCH] Draft 17 header protection --- quinn-proto/src/connection.rs | 52 ++++++++++++----------------- quinn-proto/src/crypto.rs | 61 ++++++++++++++++++++++------------- quinn-proto/src/packet.rs | 53 ++++++++++++------------------ 3 files changed, 81 insertions(+), 85 deletions(-) diff --git a/quinn-proto/src/connection.rs b/quinn-proto/src/connection.rs index 2d1424ce8..38f9e7a32 100644 --- a/quinn-proto/src/connection.rs +++ b/quinn-proto/src/connection.rs @@ -924,17 +924,11 @@ impl Connection { ecn: Option, partial_decode: PartialDecode, ) -> Option { - let mut new_crypto = None; - let crypto = if partial_decode.is_handshake() { - self.find_crypto(CryptoLevel::Initial, 0).unwrap() - } else if partial_decode.key_phase() != self.key_phase { - new_crypto = Some(self.cryptos.back().unwrap().crypto.update()); - new_crypto.as_ref().unwrap() + let header_key = if partial_decode.is_handshake() { + self.find_crypto(CryptoLevel::Initial, 0) + .unwrap() + .header_decrypt_key() } else { - // This is somewhat incorrect for pre-17 drafts, where packet number protection - // keys are supposed to get updated with key updates. Since supporting that is - // painful and will go away in draft 17, let's take the simple way out and just - // make sure we use the packet number protection key from the last 1-RTT key. let crypto_space = self.cryptos.back().unwrap(); if crypto_space.level != CryptoLevel::OneRtt { warn!( @@ -943,12 +937,12 @@ impl Connection { ); return None; } - &crypto_space.crypto + crypto_space.crypto.header_decrypt_key() }; - match partial_decode.finish(crypto.header_decrypt_key()) { + match partial_decode.finish(header_key) { Ok((packet, rest)) => { - self.handle_packet(now, ecn, packet, new_crypto); + self.handle_packet(now, ecn, packet); rest } Err(e) => { @@ -958,13 +952,7 @@ impl Connection { } } - fn handle_packet( - &mut self, - now: u64, - ecn: Option, - mut packet: Packet, - crypto_update: Option, - ) { + fn handle_packet(&mut self, now: u64, ecn: Option, mut packet: Packet) { trace!(self.log, "connection got packet"; "len" => packet.payload.len()); let was_handshake = self.is_handshaking(); let was_closed = self.state.is_closed(); @@ -973,7 +961,7 @@ impl Connection { packet.payload.len() >= 16 && packet.payload[packet.payload.len() - 16..] == token }); - let result = match self.decrypt_packet(was_handshake, &mut packet, crypto_update) { + let result = match self.decrypt_packet(was_handshake, &mut packet) { Err(Some(e)) => { warn!(self.log, "got illegal packet"; "reason" => %e); Err(e.into()) @@ -2181,16 +2169,17 @@ impl Connection { &mut self, handshake: bool, packet: &mut Packet, - crypto_update: Option, ) -> Result, Option> { if packet.header.is_retry() { // Retry packets are not encrypted and have no packet number return Ok(None); } - let (level, number) = match packet.header { - Header::Short { number, .. } if !handshake => (CryptoLevel::OneRtt, number), + let (level, number, key_phase) = match packet.header { + Header::Short { + number, key_phase, .. + } if !handshake => (CryptoLevel::OneRtt, number, key_phase), Header::Initial { number, .. } | Header::Long { number, .. } if handshake => { - (CryptoLevel::Initial, number) + (CryptoLevel::Initial, number, false) } _ => { return Err(None); @@ -2198,12 +2187,13 @@ impl Connection { }; let number = number.expand(self.rx_packet + 1); - let crypto = match crypto_update.as_ref() { - None => self.find_crypto(level, number).unwrap(), - Some(crypto) => { - assert_eq!(level, CryptoLevel::OneRtt); - crypto - } + let mut crypto_update = None; + let crypto = if key_phase == self.key_phase { + self.find_crypto(level, number).unwrap() + } else { + assert_eq!(level, CryptoLevel::OneRtt); + crypto_update = Some(self.find_crypto(level, number).unwrap().update()); + crypto_update.as_ref().unwrap() }; crypto diff --git a/quinn-proto/src/crypto.rs b/quinn-proto/src/crypto.rs index 3257bce36..dcc7b0ed3 100644 --- a/quinn-proto/src/crypto.rs +++ b/quinn-proto/src/crypto.rs @@ -20,7 +20,7 @@ pub use rustls::{ClientConfig, ClientSession, ServerConfig, ServerSession, Sessi use webpki::DNSNameRef; use crate::coding::{BufExt, BufMutExt}; -use crate::packet::{ConnectionId, AEAD_TAG_SIZE}; +use crate::packet::{ConnectionId, PacketNumber, AEAD_TAG_SIZE, LONG_HEADER_FORM}; use crate::transport_parameters::TransportParameters; use crate::{Side, MAX_CID_SIZE, MIN_CID_SIZE, RESET_TOKEN_SIZE}; @@ -356,41 +356,58 @@ impl HeaderKey { } } - pub fn decrypt(&self, sample: &[u8], in_out: &mut [u8]) { + fn mask(&self, sample: &[u8]) -> [u8; 5] { + let mut buf = [0; 5]; use self::HeaderKey::*; match self { AesCtr128(key) => { let key = GenericArray::from_slice(key); let nonce = GenericArray::from_slice(sample); - Aes128Ctr::new(key, nonce).apply_keystream(in_out) + Aes128Ctr::new(key, nonce).apply_keystream(&mut buf); } ChaCha20(key) => { let counter = BigEndian::read_u32(&sample[..4]); let nonce = chacha20::Nonce::from_slice(&sample[4..]).expect("failed to generate nonce"); - let mut input = [0; 4]; - (&mut input[..in_out.len()]).copy_from_slice(in_out); - chacha20::decrypt(key, &nonce, counter, &input[..in_out.len()], in_out).unwrap(); + chacha20::decrypt(key, &nonce, counter, &[0; 5], &mut buf).unwrap(); } } + buf + } + + pub fn decrypt(&self, pn_offset: usize, sample: &[u8], packet: &mut [u8]) { + let mask = self.mask(sample); + if packet[0] & LONG_HEADER_FORM == LONG_HEADER_FORM { + // Long header: 4 bits masked + packet[0] ^= mask[0] & 0x0f; + } else { + // Short header: 5 bits masked + packet[0] ^= mask[0] & 0x1f; + } + let pn_length = PacketNumber::decode_len(packet[0]); + for (out, inp) in packet[pn_offset..pn_offset + pn_length] + .iter_mut() + .zip(&mask[1..]) + { + *out ^= inp; + } } - pub fn encrypt(&self, sample: &[u8], in_out: &mut [u8]) { - use self::HeaderKey::*; - match self { - AesCtr128(key) => { - let key = GenericArray::from_slice(key); - let nonce = GenericArray::from_slice(sample); - Aes128Ctr::new(key, nonce).apply_keystream(in_out) - } - ChaCha20(key) => { - let counter = BigEndian::read_u32(&sample[..4]); - let nonce = - chacha20::Nonce::from_slice(&sample[4..]).expect("failed to generate nonce"); - let mut input = [0; 4]; - (&mut input[..in_out.len()]).copy_from_slice(in_out); - chacha20::encrypt(key, &nonce, counter, &input[..in_out.len()], in_out).unwrap(); - } + pub fn encrypt(&self, pn_offset: usize, sample: &[u8], packet: &mut [u8]) { + let mask = self.mask(sample); + let pn_length = PacketNumber::decode_len(packet[0]); + if packet[0] & 0x80 == 0x80 { + // Long header: 4 bits masked + packet[0] ^= mask[0] & 0x0f; + } else { + // Short header: 5 bits masked + packet[0] ^= mask[0] & 0x1f; + } + for (out, inp) in packet[pn_offset..pn_offset + pn_length] + .iter_mut() + .zip(&mask[1..]) + { + *out ^= inp; } } } diff --git a/quinn-proto/src/packet.rs b/quinn-proto/src/packet.rs index b32c97820..1a6a00905 100644 --- a/quinn-proto/src/packet.rs +++ b/quinn-proto/src/packet.rs @@ -75,13 +75,6 @@ impl PartialDecode { self.invariant_header.dst_cid() } - pub fn key_phase(&self) -> bool { - match self.invariant_header { - InvariantHeader::Short { first, .. } => (first & KEY_PHASE_BIT) != 0, - _ => false, - } - } - pub fn finish( self, header_key: &HeaderKey, @@ -91,8 +84,7 @@ impl PartialDecode { mut buf, } = self; let (payload_len, header, allow_coalesced) = match invariant_header { - InvariantHeader::Short { first, dst_cid } => { - let key_phase = first & KEY_PHASE_BIT != 0; + InvariantHeader::Short { dst_cid, .. } => { if !buf.has_remaining() { return Err(PacketDecodeError::InvalidHeader( "header ends before packet number", @@ -100,7 +92,8 @@ impl PartialDecode { } let sample_offset = 1 + dst_cid.len() + 4; - let number = Self::get_packet_number(&mut buf, header_key, sample_offset)?; + let number = Self::decrypt_header(&mut buf, header_key, sample_offset)?; + let key_phase = buf.get_ref()[0] & KEY_PHASE_BIT != 0; ( buf.remaining(), Header::Short { @@ -163,7 +156,7 @@ impl PartialDecode { + varint::size(token_length as u64).unwrap() + token.len(); - let number = Self::get_packet_number(&mut buf, header_key, sample_offset)?; + let number = Self::decrypt_header(&mut buf, header_key, sample_offset)?; ( (len as usize) - number.len(), Header::Initial { @@ -179,7 +172,7 @@ impl PartialDecode { let len = buf.get_var()?; let sample_offset = 10 + dst_cid.len() + src_cid.len() + varint::size(len).unwrap(); - let number = Self::get_packet_number(&mut buf, header_key, sample_offset)?; + let number = Self::decrypt_header(&mut buf, header_key, sample_offset)?; ( (len as usize) - number.len(), Header::Long { @@ -218,7 +211,7 @@ impl PartialDecode { )) } - fn get_packet_number( + fn decrypt_header( buf: &mut io::Cursor, header_key: &HeaderKey, mut sample_offset: usize, @@ -242,11 +235,10 @@ impl PartialDecode { sample.copy_from_slice( &buf.get_ref()[sample_offset..sample_offset + header_key.sample_size()], ); + let pn_offset = buf.position() as usize; + header_key.decrypt(pn_offset, &sample, buf.get_mut()); - let pos = buf.position() as usize; let len = PacketNumber::decode_len(buf.get_ref()[0]); - - header_key.decrypt(&sample, &mut buf.get_mut()[pos..pos + len]); PacketNumber::decode(len, buf) } } @@ -312,7 +304,7 @@ impl Header { + token.len(); PartialEncode { header: self, - pn: Some((pn_pos, number.len())), + pn: Some(pn_pos), } } Long { @@ -329,7 +321,7 @@ impl Header { let pn_pos = 8 + dst_cid.len() + src_cid.len(); PartialEncode { header: self, - pn: Some((pn_pos, number.len())), + pn: Some(pn_pos), } } Retry { @@ -357,7 +349,7 @@ impl Header { number.encode(w); PartialEncode { header: self, - pn: Some((1 + dst_cid.len(), number.len())), + pn: Some(1 + dst_cid.len()), } } VersionNegotiate { @@ -400,18 +392,17 @@ impl Header { pub struct PartialEncode<'a> { header: &'a Header, - pn: Option<(usize, usize)>, + pn: Option, } impl<'a> PartialEncode<'a> { pub fn finish(self, buf: &mut [u8], header_key: &HeaderKey, header_len: usize) { let PartialEncode { header, pn } = self; let payload_len = (buf.len() - header_len) as u64; - let (mut sample_offset, pn_pos, pn_len) = match header { + let (mut sample_offset, pn_pos) = match header { Header::Short { dst_cid, .. } => { let sample_offset = 1 + dst_cid.len() + 4; - let (pn_pos, pn_len) = pn.unwrap(); - (sample_offset, pn_pos, pn_len) + (sample_offset, pn.unwrap()) } Header::Initial { dst_cid, @@ -425,16 +416,14 @@ impl<'a> PartialEncode<'a> { + varint::size(payload_len).unwrap() + varint::size(token.len() as u64).unwrap() + token.len(); - let (pn_pos, pn_len) = pn.unwrap(); - (sample_offset, pn_pos, pn_len) + (sample_offset, pn.unwrap()) } Header::Long { dst_cid, src_cid, .. } => { let sample_offset = 10 + dst_cid.len() + src_cid.len() + varint::size(payload_len).unwrap(); - let (pn_pos, pn_len) = pn.unwrap(); - (sample_offset, pn_pos, pn_len) + (sample_offset, pn.unwrap()) } _ => { return; @@ -453,7 +442,7 @@ impl<'a> PartialEncode<'a> { sample }; - header_key.encrypt(&sample, &mut buf[pn_pos..pn_pos + pn_len]); + header_key.encrypt(pn_pos, &sample, buf); } } @@ -589,7 +578,7 @@ impl PacketNumber { Ok(pn) } - fn decode_len(tag: u8) -> usize { + pub fn decode_len(tag: u8) -> usize { 1 + (tag & 0x03) as usize } @@ -808,7 +797,7 @@ pub fn set_payload_length(packet: &mut [u8], header_len: usize, pn_len: usize) { pub const AEAD_TAG_SIZE: usize = 16; -const LONG_HEADER_FORM: u8 = 0x80; +pub const LONG_HEADER_FORM: u8 = 0x80; const KEY_PHASE_BIT: u8 = 0x04; const FIXED_BIT: u8 = 0x40; @@ -914,7 +903,7 @@ mod tests { }; PartialEncode { header: &header, - pn: Some((1, 2)), + pn: Some(1), } .finish(&mut sending, &key, 3); assert_eq!(&sending[1..3], [0x80, 0x6d]); @@ -959,7 +948,7 @@ mod tests { }; PartialEncode { header: &header, - pn: Some((1, 2)), + pn: Some(1), } .finish(&mut sending, &key, 3); assert_eq!(&sending[1..3], [0xa9, 0x0e]);