Draft 17 header protection

This commit is contained in:
Benjamin Saunders
2018-12-23 14:09:35 -08:00
committed by Dirkjan Ochtman
parent d408867d97
commit 28aa2f1f9b
3 changed files with 81 additions and 85 deletions
+21 -31
View File
@@ -924,17 +924,11 @@ impl Connection {
ecn: Option<EcnCodepoint>,
partial_decode: PartialDecode,
) -> Option<BytesMut> {
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<EcnCodepoint>,
mut packet: Packet,
crypto_update: Option<Crypto>,
) {
fn handle_packet(&mut self, now: u64, ecn: Option<EcnCodepoint>, 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<Crypto>,
) -> Result<Option<u64>, Option<TransportError>> {
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
+39 -22
View File
@@ -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;
}
}
}
+21 -32
View File
@@ -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<BytesMut>,
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<usize>,
}
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]);