mirror of
https://github.com/n0-computer/noq.git
synced 2026-09-21 02:33:23 +00:00
Draft 17 header protection
This commit is contained in:
committed by
Dirkjan Ochtman
parent
d408867d97
commit
28aa2f1f9b
@@ -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
@@ -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
@@ -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]);
|
||||
|
||||
Reference in New Issue
Block a user