From 2c21532ea6ee4c7381d583e552fbd358f7e884ea Mon Sep 17 00:00:00 2001 From: Dirkjan Ochtman Date: Mon, 12 Nov 2018 17:27:04 +0100 Subject: [PATCH] Apply packet number encryption and decryption --- quinn-proto/src/connection.rs | 46 ++++++--- quinn-proto/src/endpoint.rs | 2 +- quinn-proto/src/packet.rs | 188 +++++++++++++++++++++++++++++++--- 3 files changed, 208 insertions(+), 28 deletions(-) diff --git a/quinn-proto/src/connection.rs b/quinn-proto/src/connection.rs index d4d818b97..9ef151479 100644 --- a/quinn-proto/src/connection.rs +++ b/quinn-proto/src/connection.rs @@ -1057,7 +1057,16 @@ impl Connection { remote: SocketAddrV6, partial_decode: PartialDecode, ) -> Option { - match partial_decode.finish() { + let result = { + let crypto = if partial_decode.is_handshake() { + &self.handshake_crypto + } else { + &self.crypto.as_ref().unwrap() + }; + partial_decode.finish(crypto.pn_decrypt_key()) + }; + + match result { Ok((packet, rest)) => { self.handle_packet(ctx, now, remote, packet); rest @@ -1966,7 +1975,7 @@ impl Connection { return None; }; - header.encode(&mut buf); + let partial_encode = header.encode(&mut buf); let ack_only = pending.is_empty(); let header_len = buf.len() as u16; let max_size = self.mtu as usize - AEAD_TAG_SIZE; @@ -2156,6 +2165,7 @@ impl Connection { set_payload_length(&mut buf, header_len as usize, pn_len); } crypto.encrypt(number, &mut buf, header_len as usize); + partial_encode.finish(&mut buf, crypto.pn_encrypt_key(), header_len as usize); (number, acks, ack_only, crypto.is_initial()) }; @@ -2184,17 +2194,19 @@ impl Connection { fn force_transmit(&mut self, config: &Config, now: u64) -> Box<[u8]> { let number = self.get_tx_number(); let mut buf = Vec::new(); - Header::Short { + let header = Header::Short { dst_cid: self.rem_cid, number: PacketNumber::new(number, self.largest_acked_packet), key_phase: self.key_phase, - }.encode(&mut buf); + }; + let partial_encode = header.encode(&mut buf); let header_len = buf.len() as u16; buf.push(frame::Type::PING.into()); - self.crypto - .as_ref() - .unwrap() - .encrypt(number, &mut buf, header_len as usize); + { + let crypto = self.crypto.as_ref().unwrap(); + crypto.encrypt(number, &mut buf, header_len as usize); + partial_encode.finish(&mut buf, crypto.pn_encrypt_key(), header_len as usize); + } self.on_packet_sent( config, now, @@ -2213,21 +2225,26 @@ impl Connection { fn make_close(&mut self, reason: &state::CloseReason) -> Box<[u8]> { let number = self.get_tx_number(); let mut buf = Vec::new(); - Header::Short { + let header = Header::Short { dst_cid: self.rem_cid, number: PacketNumber::new(number, self.largest_acked_packet), key_phase: self.key_phase, - }.encode(&mut buf); + }; + let partial_encode = header.encode(&mut buf); let header_len = buf.len() as u16; + let max_len = self.mtu - header_len - AEAD_TAG_SIZE as u16; match *reason { state::CloseReason::Application(ref x) => x.encode(&mut buf, max_len), state::CloseReason::Connection(ref x) => x.encode(&mut buf, max_len), } - self.crypto + + let crypto = self + .crypto .as_ref() - .unwrap_or_else(|| &self.handshake_crypto) - .encrypt(number, &mut buf, header_len as usize); + .unwrap_or_else(|| &self.handshake_crypto); + crypto.encrypt(number, &mut buf, header_len as usize); + partial_encode.finish(&mut buf, crypto.pn_encrypt_key(), header_len as usize); buf.into() } @@ -2698,7 +2715,7 @@ where }; let mut buf = Vec::::new(); - header.encode(&mut buf); + let partial_encode = header.encode(&mut buf); let header_len = buf.len(); let max_len = MIN_MTU - header_len as u16 - AEAD_TAG_SIZE as u16; match reason.into() { @@ -2717,6 +2734,7 @@ where } set_payload_length(&mut buf, header_len, number.len()); crypto.encrypt(packet_number as u64, &mut buf, header_len); + partial_encode.finish(&mut buf, crypto.pn_encrypt_key(), header_len); buf.into() } diff --git a/quinn-proto/src/endpoint.rs b/quinn-proto/src/endpoint.rs index abfa3a480..92f13dc21 100644 --- a/quinn-proto/src/endpoint.rs +++ b/quinn-proto/src/endpoint.rs @@ -373,7 +373,7 @@ impl Endpoint { } let crypto = Crypto::new_initial(&partial_decode.dst_cid(), Side::Server); - return match partial_decode.finish() { + return match partial_decode.finish(crypto.pn_decrypt_key()) { Ok((packet, rest)) => { self.handle_initial(now, remote, packet, crypto); rest diff --git a/quinn-proto/src/packet.rs b/quinn-proto/src/packet.rs index 045baa63b..8bd24e2d3 100644 --- a/quinn-proto/src/packet.rs +++ b/quinn-proto/src/packet.rs @@ -5,6 +5,8 @@ use rand::Rng; use slog; use coding::{self, BufExt, BufMutExt, Codec}; +use crypto::PacketNumberKey; +use varint; use {MAX_CID_SIZE, MIN_CID_SIZE, VERSION}; // Due to packet number encryption, it is impossible to fully decode a header @@ -52,11 +54,30 @@ impl PartialDecode { } } + pub fn is_handshake(&self) -> bool { + match self.invariant_header { + InvariantHeader::Long { + version: VERSION, + first, + .. + } => match PacketType::from_byte(first).unwrap() { + PacketType::Initial => true, + PacketType::Long(LongType::Handshake) => true, + _ => false, + }, + InvariantHeader::Long { .. } => false, + InvariantHeader::Short { .. } => false, + } + } + pub fn dst_cid(&self) -> ConnectionId { self.invariant_header.dst_cid() } - pub fn finish(self) -> Result<(Packet, Option), PacketDecodeError> { + pub fn finish( + self, + pn_key: &PacketNumberKey, + ) -> Result<(Packet, Option), PacketDecodeError> { let Self { invariant_header, mut buf, @@ -64,7 +85,15 @@ impl PartialDecode { let (payload_len, header, allow_coalesced) = match invariant_header { InvariantHeader::Short { first, dst_cid } => { let key_phase = first & KEY_PHASE_BIT != 0; - let number = PacketNumber::decode(&mut buf)?; + if !buf.has_remaining() { + return Err(PacketDecodeError::InvalidHeader( + "header ends before packet number", + )); + } + + let mut sample_offset = 1 + dst_cid.len() + 4; + let number = Self::get_packet_number(&mut buf, pn_key, sample_offset)?; + ( buf.remaining(), Header::Short { @@ -117,7 +146,14 @@ impl PartialDecode { buf.copy_to_slice(&mut token); let len = buf.get_var()?; - let number = PacketNumber::decode(&mut buf)?; + let sample_offset = 10 + + dst_cid.len() + + src_cid.len() + + varint::size(len).unwrap() + + varint::size(token_length as u64).unwrap() + + token.len(); + + let number = Self::get_packet_number(&mut buf, pn_key, sample_offset)?; ( len as usize, Header::Initial { @@ -131,7 +167,9 @@ impl PartialDecode { } PacketType::Long(ty) => { let len = buf.get_var()?; - let number = PacketNumber::decode(&mut buf)?; + let sample_offset = + 10 + dst_cid.len() + src_cid.len() + varint::size(len).unwrap(); + let number = Self::get_packet_number(&mut buf, pn_key, sample_offset)?; ( len as usize, Header::Long { @@ -169,6 +207,42 @@ impl PartialDecode { if allow_coalesced { Some(bytes) } else { None }, )) } + + fn get_packet_number( + buf: &mut io::Cursor, + pn_key: &PacketNumberKey, + mut sample_offset: usize, + ) -> Result { + let packet_length = buf.get_ref().len(); + if sample_offset + pn_key.sample_size() > packet_length { + sample_offset = packet_length + .checked_sub(pn_key.sample_size()) + .ok_or_else(|| { + PacketDecodeError::InvalidHeader("packet too short to decode packet number") + })?; + } + if packet_length < sample_offset + pn_key.sample_size() { + return Err(PacketDecodeError::InvalidHeader( + "packet too short to extract packet number encryption sample", + )); + } + + let mut first = [buf.bytes()[0]; 1]; + let sample = { + let mut sample = [0; 16]; + debug_assert!(pn_key.sample_size() <= 16); + sample.copy_from_slice( + &buf.get_ref()[sample_offset..sample_offset + pn_key.sample_size()], + ); + sample + }; + + pn_key.decrypt(&sample, &mut first); + let len = PacketNumber::decode_len(first[0]); + let pos = buf.position() as usize; + pn_key.decrypt(&sample, &mut buf.get_mut()[pos..pos + len]); + PacketNumber::decode(buf) + } } pub struct Packet { @@ -209,7 +283,7 @@ pub enum Header { } impl Header { - pub fn encode(&self, w: &mut W) { + pub fn encode(&self, w: &mut W) -> PartialEncode { use self::Header::*; match *self { Initial { @@ -225,6 +299,15 @@ impl Header { w.put_slice(token); w.write::(0); // Placeholder for payload length; see `set_payload_length` number.encode(w); + let pn_pos = 8 + + dst_cid.len() + + src_cid.len() + + varint::size(token.len() as u64).unwrap() + + token.len(); + PartialEncode { + header: self, + pn: Some((pn_pos, number.len())), + } } Long { ty, @@ -237,6 +320,11 @@ impl Header { Self::encode_cids(w, dst_cid, src_cid); w.write::(0); // Placeholder for payload length; see `set_payload_length` number.encode(w); + let pn_pos = 8 + dst_cid.len() + src_cid.len(); + PartialEncode { + header: self, + pn: Some((pn_pos, number.len())), + } } Retry { ref src_cid, @@ -248,6 +336,10 @@ impl Header { Self::encode_cids(w, dst_cid, src_cid); w.write(orig_dst_cid.len() as u8); w.put_slice(orig_dst_cid); + PartialEncode { + header: self, + pn: None, + } } Short { ref dst_cid, @@ -257,6 +349,10 @@ impl Header { w.write(0x30 | if key_phase { KEY_PHASE_BIT } else { 0 }); w.put_slice(dst_cid); number.encode(w); + PartialEncode { + header: self, + pn: Some((1 + dst_cid.len(), number.len())), + } } VersionNegotiate { ref random, @@ -266,6 +362,10 @@ impl Header { w.write(0x80u8 | random); w.write::(0); Self::encode_cids(w, dst_cid, src_cid); + PartialEncode { + header: self, + pn: None, + } } } } @@ -285,6 +385,65 @@ impl Header { } } +pub struct PartialEncode<'a> { + header: &'a Header, + pn: Option<(usize, usize)>, +} + +impl<'a> PartialEncode<'a> { + pub fn finish(self, buf: &mut [u8], pn_key: &PacketNumberKey, 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 { + Header::Short { dst_cid, .. } => { + let mut sample_offset = 1 + dst_cid.len() + 4; + let (pn_pos, pn_len) = pn.unwrap(); + (sample_offset, pn_pos, pn_len) + } + Header::Initial { + dst_cid, + src_cid, + token, + .. + } => { + let sample_offset = 10 + + dst_cid.len() + + src_cid.len() + + 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) + } + 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) + } + _ => { + return; + } + }; + + let packet_length = buf.len(); + if sample_offset + pn_key.sample_size() > packet_length { + sample_offset = packet_length - pn_key.sample_size(); + } + + debug_assert!(pn_key.sample_size() <= 16); + let sample = { + let mut sample = [0; 16]; + sample.copy_from_slice(&buf[sample_offset..sample_offset + pn_key.sample_size()]); + sample + }; + + pn_key.encrypt(&sample, &mut buf[pn_pos..pn_pos + pn_len]); + } +} + pub enum InvariantHeader { Long { first: u8, @@ -417,14 +576,7 @@ impl PacketNumber { return Err(coding::UnexpectedEnd.into()); } - let first = r.bytes()[0]; - let len = if first < 0x80 { - 1 - } else if first < 0xc0 { - 2 - } else { - 4 - }; + let len = Self::decode_len(r.bytes()[0]); if r.remaining() < len { return Err(coding::UnexpectedEnd.into()); } @@ -439,6 +591,16 @@ impl PacketNumber { } } + fn decode_len(b: u8) -> usize { + if b < 0x80 { + 1 + } else if b < 0xc0 { + 2 + } else { + 4 + } + } + pub fn expand(self, prev: u64) -> u64 { use self::PacketNumber::*; let t = prev + 1;