Apply packet number encryption and decryption

This commit is contained in:
Dirkjan Ochtman
2018-11-12 17:27:04 +01:00
parent 4342c92712
commit 2c21532ea6
3 changed files with 208 additions and 28 deletions
+32 -14
View File
@@ -1057,7 +1057,16 @@ impl Connection {
remote: SocketAddrV6,
partial_decode: PartialDecode,
) -> Option<BytesMut> {
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::<u8>::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()
}
+1 -1
View File
@@ -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
+175 -13
View File
@@ -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<BytesMut>), PacketDecodeError> {
pub fn finish(
self,
pn_key: &PacketNumberKey,
) -> Result<(Packet, Option<BytesMut>), 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<BytesMut>,
pn_key: &PacketNumberKey,
mut sample_offset: usize,
) -> Result<PacketNumber, PacketDecodeError> {
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<W: BufMut>(&self, w: &mut W) {
pub fn encode<W: BufMut>(&self, w: &mut W) -> PartialEncode {
use self::Header::*;
match *self {
Initial {
@@ -225,6 +299,15 @@ impl Header {
w.put_slice(token);
w.write::<u16>(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::<u16>(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::<u32>(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;