use core::{mem::size_of, time::Duration};
use s2n_codec::{DecoderBuffer, DecoderBufferMut};
use s2n_quic_core::{
connection, event::api::SocketAddress, random, time::Timestamp, token::Source,
};
use s2n_quic_crypto::{constant_time, digest, hmac};
use std::collections::HashSet;
use zerocopy::{FromBytes, IntoBytes, Unaligned};
use zeroize::Zeroizing;
#[derive(Debug, Default)]
struct DuplicateFilter {
seen: HashSet<[u8; 32]>,
}
impl DuplicateFilter {
const MAX_ENTRIES: usize = 16 * 1024;
fn insert(&mut self, token: &Token) -> bool {
if self.seen.len() < Self::MAX_ENTRIES {
self.seen.insert(token.hmac)
} else {
false
}
}
}
struct BaseKey {
active_duration: Duration,
key: Option<(Timestamp, hmac::Key)>,
duplicate_filter: DuplicateFilter,
}
impl BaseKey {
pub fn new(active_duration: Duration) -> Self {
Self {
active_duration,
key: None,
duplicate_filter: DuplicateFilter::default(),
}
}
pub fn hasher(&mut self, random: &mut dyn random::Generator) -> Option<hmac::Context> {
let key = self.poll_key(random)?;
Some(hmac::Context::with_key(&key))
}
fn poll_key(&mut self, random: &mut dyn random::Generator) -> Option<hmac::Key> {
let now = s2n_quic_platform::time::now();
if let Some((expires_at, key)) = self.key.as_ref() {
if expires_at > &now {
return Some(key.clone());
}
}
let expires_at = now.checked_add(self.active_duration)?;
let mut key_material = Zeroizing::new([0; digest::SHA256_OUTPUT_LEN]);
random.private_random_fill(&mut key_material[..]);
let key = hmac::Key::new(hmac::HMAC_SHA256, key_material.as_ref());
self.duplicate_filter = DuplicateFilter::default();
self.key = Some((expires_at, key));
self.key.as_ref().map(|key| key.1.clone())
}
}
const DEFAULT_KEY_ROTATION_PERIOD: Duration = Duration::from_millis(1000);
#[derive(Debug)]
pub struct Provider {
key_rotation_period: Duration,
}
impl Default for Provider {
fn default() -> Self {
Self {
key_rotation_period: DEFAULT_KEY_ROTATION_PERIOD,
}
}
}
impl super::Provider for Provider {
type Format = Format;
type Error = core::convert::Infallible;
fn start(self) -> Result<Self::Format, Self::Error> {
let format = Format {
key_rotation_period: self.key_rotation_period,
current_key_rotates_at: s2n_quic_platform::time::now(),
current_key: 0,
keys: [
BaseKey::new(self.key_rotation_period * 2),
BaseKey::new(self.key_rotation_period * 2),
],
};
Ok(format)
}
}
pub struct Format {
key_rotation_period: Duration,
current_key_rotates_at: s2n_quic_core::time::Timestamp,
current_key: u8,
keys: [BaseKey; 2],
}
impl Format {
fn current_key(&mut self) -> u8 {
let now = s2n_quic_platform::time::now();
if now > self.current_key_rotates_at {
self.current_key ^= 1;
self.current_key_rotates_at = now + self.key_rotation_period;
}
self.current_key
}
fn tag_retry_token(
&mut self,
token: &Token,
context: &mut super::Context<'_>,
) -> Option<hmac::Tag> {
let mut ctx = self.keys[token.header.key_id() as usize].hasher(context.random)?;
ctx.update(&token.original_destination_connection_id);
ctx.update(&token.nonce);
ctx.update(context.peer_connection_id);
match context.remote_address {
SocketAddress::IpV4 { ip, port, .. } => {
ctx.update(ip);
ctx.update(&port.to_be_bytes());
}
SocketAddress::IpV6 { ip, port, .. } => {
ctx.update(ip);
ctx.update(&port.to_be_bytes());
}
_ => {
return None;
}
};
Some(ctx.sign())
}
fn validate_retry_token(
&mut self,
context: &mut super::Context<'_>,
token: &Token,
) -> Option<connection::InitialId> {
let tag = self.tag_retry_token(token, context)?;
if constant_time::verify_slices_are_equal(&token.hmac, tag.as_ref()).is_err() {
return None;
}
if !self.keys[token.header.key_id() as usize]
.duplicate_filter
.insert(token)
{
return None;
}
token.original_destination_connection_id()
}
}
impl super::Format for Format {
const TOKEN_LEN: usize = size_of::<Token>();
fn generate_new_token(
&mut self,
_context: &mut super::Context<'_>,
_source_connection_id: &connection::LocalId,
_output_buffer: &mut [u8],
) -> Option<()> {
None
}
fn generate_retry_token(
&mut self,
context: &mut super::Context<'_>,
original_destination_connection_id: &connection::InitialId,
output_buffer: &mut [u8],
) -> Option<()> {
let buffer = DecoderBufferMut::new(output_buffer);
let (token, _) = buffer
.decode::<&mut Token>()
.expect("Provided output buffer did not match TOKEN_LEN");
let header = Header::new(Source::RetryPacket, self.current_key());
token.header = header;
token.original_destination_connection_id[..original_destination_connection_id.len()]
.copy_from_slice(original_destination_connection_id.as_bytes());
token.odcid_len = original_destination_connection_id.len() as u8;
for b in token
.original_destination_connection_id
.iter_mut()
.skip(original_destination_connection_id.len())
{
*b = 0;
}
context.random.public_random_fill(&mut token.nonce[..]);
let tag = self.tag_retry_token(token, context)?;
token.hmac.copy_from_slice(tag.as_ref());
Some(())
}
fn validate_token(
&mut self,
context: &mut super::Context<'_>,
token: &[u8],
) -> Option<connection::InitialId> {
let buffer = DecoderBuffer::new(token);
let (token, remaining) = buffer.decode::<&Token>().ok()?;
remaining.ensure_empty().ok()?;
if token.header.version() != TOKEN_VERSION {
return None;
}
let source = token.header.token_source();
match source {
Source::RetryPacket => self.validate_retry_token(context, token),
Source::NewTokenFrame => None, }
}
}
#[derive(Clone, Copy, Debug, FromBytes, IntoBytes, Unaligned)]
#[repr(C)]
pub(crate) struct Header(u8);
const TOKEN_VERSION: u8 = 0x00;
const VERSION_SHIFT: u8 = 7;
const VERSION_MASK: u8 = 0x80;
const TOKEN_SOURCE_SHIFT: u8 = 6;
const TOKEN_SOURCE_MASK: u8 = 0x40;
const KEY_ID_SHIFT: u8 = 5;
const KEY_ID_MASK: u8 = 0x20;
impl Header {
fn new(source: Source, key_id: u8) -> Header {
let mut header: u8 = 0;
header |= TOKEN_VERSION << VERSION_SHIFT;
header |= match source {
Source::NewTokenFrame => 0 << TOKEN_SOURCE_SHIFT,
Source::RetryPacket => 1 << TOKEN_SOURCE_SHIFT,
};
debug_assert!(key_id <= 1);
header |= (key_id & 0x01) << KEY_ID_SHIFT;
Header(header)
}
fn version(self) -> u8 {
(self.0 & VERSION_MASK) >> VERSION_SHIFT
}
fn key_id(self) -> u8 {
(self.0 & KEY_ID_MASK) >> KEY_ID_SHIFT
}
fn token_source(self) -> Source {
match (self.0 & TOKEN_SOURCE_MASK) >> TOKEN_SOURCE_SHIFT {
0 => Source::NewTokenFrame,
1 => Source::RetryPacket,
_ => Source::NewTokenFrame,
}
}
}
#[derive(Copy, Clone, Debug, FromBytes, IntoBytes, Unaligned)]
#[repr(C)]
struct Token {
header: Header,
odcid_len: u8,
original_destination_connection_id: [u8; 20],
nonce: [u8; 32],
hmac: [u8; 32],
}
s2n_codec::zerocopy_value_codec!(Token);
impl Token {
pub fn original_destination_connection_id(&self) -> Option<connection::InitialId> {
let dcid = self
.original_destination_connection_id
.get(..self.odcid_len as usize)?;
connection::InitialId::try_from_bytes(dcid)
}
}
#[cfg(test)]
mod tests {
use super::*;
use s2n_quic_core::{
inet::SocketAddress,
random,
token::{Context, Format as FormatTrait, Source},
};
use s2n_quic_platform::time;
use std::{net::SocketAddr, sync::Arc};
use zerocopy::FromZeros;
const TEST_KEY_ROTATION_PERIOD: Duration = Duration::from_millis(1000);
fn get_test_format() -> Format {
Format {
key_rotation_period: TEST_KEY_ROTATION_PERIOD,
keys: [
BaseKey::new(TEST_KEY_ROTATION_PERIOD * 2),
BaseKey::new(TEST_KEY_ROTATION_PERIOD * 2),
],
current_key_rotates_at: time::now(),
current_key: 0,
}
}
#[test]
fn test_header() {
for source in &[Source::NewTokenFrame, Source::RetryPacket] {
for key_id in [0, 1] {
let header = Header::new(*source, key_id);
assert_eq!(header.version(), TOKEN_VERSION);
assert_eq!(header.token_source(), *source);
assert_eq!(header.key_id(), key_id);
}
}
}
#[test]
fn test_valid_retry_tokens() {
let clock = Arc::new(time::testing::MockClock::new());
time::testing::set_local_clock(clock.clone());
let mut format = get_test_format();
let first_conn_id = connection::PeerId::try_from_bytes(&[2, 4, 6, 8, 10]).unwrap();
let second_conn_id = connection::PeerId::try_from_bytes(&[1, 3, 5, 7, 9]).unwrap();
let orig_conn_id =
connection::InitialId::try_from_bytes(&[0, 1, 2, 3, 4, 5, 6, 7]).unwrap();
let addr = SocketAddress::default();
let mut first_token = [0; Format::TOKEN_LEN];
let mut second_token = [0; Format::TOKEN_LEN];
let mut random = random::testing::Generator(5);
let mut context = Context::new(&addr, &first_conn_id, &mut random);
format
.generate_retry_token(&mut context, &orig_conn_id, &mut first_token)
.unwrap();
context = Context::new(&addr, &second_conn_id, &mut random);
format
.generate_retry_token(&mut context, &orig_conn_id, &mut second_token)
.unwrap();
clock.adjust_by(TEST_KEY_ROTATION_PERIOD);
context = Context::new(&addr, &first_conn_id, &mut random);
assert_eq!(
format.validate_token(&mut context, &first_token),
Some(orig_conn_id)
);
context = Context::new(&addr, &second_conn_id, &mut random);
assert_eq!(
format.validate_token(&mut context, &second_token),
Some(orig_conn_id)
);
context = Context::new(&addr, &first_conn_id, &mut random);
assert_eq!(format.validate_token(&mut context, &second_token), None);
}
#[test]
fn test_retry_ip_port_validation() {
let mut format = get_test_format();
let conn_id = connection::PeerId::try_from_bytes(&[2, 4, 6, 8, 10]).unwrap();
let orig_conn_id =
connection::InitialId::try_from_bytes(&[0, 1, 2, 3, 4, 5, 6, 7]).unwrap();
let mut token = [0; Format::TOKEN_LEN];
let ip_address = "127.0.0.1:443";
let addr: SocketAddr = ip_address.parse().unwrap();
let correct_address: SocketAddress = addr.into();
let mut random = random::testing::Generator(5);
let mut context = Context::new(&correct_address, &conn_id, &mut random);
format
.generate_retry_token(&mut context, &orig_conn_id, &mut token)
.unwrap();
let ip_address = "127.0.0.2:443";
let addr: SocketAddr = ip_address.parse().unwrap();
let incorrect_address: SocketAddress = addr.into();
context = Context::new(&incorrect_address, &conn_id, &mut random);
assert_eq!(format.validate_token(&mut context, &token), None);
let ip_address = "127.0.0.1:444";
let addr: SocketAddr = ip_address.parse().unwrap();
let incorrect_port: SocketAddress = addr.into();
context = Context::new(&incorrect_port, &conn_id, &mut random);
assert_eq!(format.validate_token(&mut context, &token), None);
context = Context::new(&correct_address, &conn_id, &mut random);
assert!(format.validate_token(&mut context, &token).is_some());
}
#[test]
fn test_key_rotation() {
let clock = Arc::new(time::testing::MockClock::new());
time::testing::set_local_clock(clock.clone());
let mut format = get_test_format();
let conn_id = connection::PeerId::TEST_ID;
let orig_conn_id = connection::InitialId::TEST_ID;
let addr = SocketAddress::default();
let mut buf = [0; Format::TOKEN_LEN];
let mut random = random::testing::Generator(5);
let mut context = Context::new(&addr, &conn_id, &mut random);
format
.generate_retry_token(&mut context, &orig_conn_id, &mut buf)
.unwrap();
clock.adjust_by(TEST_KEY_ROTATION_PERIOD);
assert!(format.validate_token(&mut context, &buf).is_some());
clock.adjust_by(TEST_KEY_ROTATION_PERIOD);
assert!(format.validate_token(&mut context, &buf).is_none());
}
#[test]
fn test_expired_retry_token() {
let clock = Arc::new(time::testing::MockClock::new());
time::testing::set_local_clock(clock.clone());
let mut format = get_test_format();
let conn_id = connection::PeerId::TEST_ID;
let orig_conn_id = connection::InitialId::TEST_ID;
let addr = SocketAddress::default();
let mut buf = [0; Format::TOKEN_LEN];
let mut random = random::testing::Generator(5);
let mut context = Context::new(&addr, &conn_id, &mut random);
format
.generate_retry_token(&mut context, &orig_conn_id, &mut buf)
.unwrap();
clock.adjust_by(TEST_KEY_ROTATION_PERIOD * 2);
assert!(format.validate_token(&mut context, &buf).is_none());
}
#[test]
fn test_retry_validation_default_format() {
let clock = Arc::new(time::testing::MockClock::new());
time::testing::set_local_clock(clock);
let mut format = get_test_format();
let conn_id = connection::PeerId::TEST_ID;
let odcid = connection::InitialId::try_from_bytes(&[0, 1, 2, 3, 4, 5, 6, 7]).unwrap();
let addr = SocketAddress::default();
let mut buf = [0; Format::TOKEN_LEN];
let mut random = random::testing::Generator(5);
let mut context = Context::new(&addr, &conn_id, &mut random);
format
.generate_retry_token(&mut context, &odcid, &mut buf)
.unwrap();
assert_eq!(format.validate_token(&mut context, &buf), Some(odcid));
let wrong_conn_id = connection::PeerId::try_from_bytes(&[0, 1, 2]).unwrap();
context = Context::new(&addr, &wrong_conn_id, &mut random);
assert!(format.validate_token(&mut context, &buf).is_none());
}
#[test]
fn test_duplicate_token_detection() {
let mut format = get_test_format();
let conn_id = connection::PeerId::TEST_ID;
let odcid = connection::InitialId::try_from_bytes(&[0, 1, 2, 3, 4, 5, 6, 7]).unwrap();
let addr = SocketAddress::default();
let mut buf = [0; Format::TOKEN_LEN];
let mut random = random::testing::Generator(5);
let mut context = Context::new(&addr, &conn_id, &mut random);
format
.generate_retry_token(&mut context, &odcid, &mut buf)
.unwrap();
assert_eq!(format.validate_token(&mut context, &buf), Some(odcid));
assert!(format.validate_token(&mut context, &buf).is_none());
}
fn test_token_with_hmac(byte: u8) -> Token {
let mut token = Token::new_zeroed();
token.hmac = [byte; 32];
token
}
#[test]
fn duplicate_filter_tracks_seen_tokens() {
let mut filter = DuplicateFilter::default();
let token = test_token_with_hmac(1);
assert!(filter.insert(&token));
assert!(!filter.insert(&token));
let other = test_token_with_hmac(2);
assert!(filter.insert(&other));
}
#[test]
fn duplicate_filter_bounds_capacity() {
let mut filter = DuplicateFilter::default();
for i in 0..DuplicateFilter::MAX_ENTRIES {
let mut token = Token::new_zeroed();
token.hmac[..8].copy_from_slice(&(i as u64).to_be_bytes());
assert!(filter.insert(&token));
}
assert_eq!(filter.seen.len(), DuplicateFilter::MAX_ENTRIES);
let overflow = test_token_with_hmac(0xff);
assert!(!filter.insert(&overflow));
assert_eq!(filter.seen.len(), DuplicateFilter::MAX_ENTRIES);
assert!(!filter.seen.contains(&overflow.hmac));
let mut seen = Token::new_zeroed();
seen.hmac[..8].copy_from_slice(&0u64.to_be_bytes());
assert!(filter.seen.contains(&seen.hmac));
assert!(!filter.insert(&seen));
}
#[test]
fn test_token_modification_detection() {
let mut format = get_test_format();
let conn_id = connection::PeerId::try_from_bytes(&[2, 4, 6, 8, 10]).unwrap();
let orig_conn_id =
connection::InitialId::try_from_bytes(&[0, 1, 2, 3, 4, 5, 6, 7]).unwrap();
let addr = SocketAddress::default();
let mut token = [0; Format::TOKEN_LEN];
let mut random = random::testing::Generator(5);
let mut context = Context::new(&addr, &conn_id, &mut random);
format
.generate_retry_token(&mut context, &orig_conn_id, &mut token)
.unwrap();
for i in 0..Format::TOKEN_LEN {
random = random::testing::Generator(5);
context = Context::new(&addr, &conn_id, &mut random);
token[i] = !token[i];
assert!(format.validate_token(&mut context, &token).is_none());
token[i] = !token[i];
}
}
#[test]
fn test_token_length_check() {
let mut format = get_test_format();
let conn_id = connection::PeerId::try_from_bytes(&[2, 4, 6, 8, 10]).unwrap();
let addr = SocketAddress::default();
bolero::check!().for_each(move |token| {
let mut random = random::testing::Generator(5);
let mut context = Context::new(&addr, &conn_id, &mut random);
assert!(format.validate_token(&mut context, token).is_none())
});
}
#[test]
fn test_token_falsification_detection() {
let mut format = get_test_format();
let conn_id = connection::PeerId::try_from_bytes(&[2, 4, 6, 8, 10]).unwrap();
let addr = SocketAddress::default();
let generator = bolero::generator::produce::<Vec<u8>>()
.with()
.len(Format::TOKEN_LEN);
bolero::check!()
.with_generator(generator)
.for_each(move |token| {
let mut random = random::testing::Generator(5);
let mut context = Context::new(&addr, &conn_id, &mut random);
assert!(format.validate_token(&mut context, token).is_none())
});
}
}