use super::{
open, receiver,
schedule::{self, Initiator},
seal, sender, stateless_reset,
};
use crate::{
credentials::{Credentials, Id},
crypto,
packet::{secret_control as control, Packet, WireVersion},
stream::TransportFeatures,
};
use rand::Rng as _;
use s2n_codec::EncoderBuffer;
use s2n_quic_core::{
dc::{self, ApplicationParams, DatagramInfo},
ensure,
event::api::EndpointType,
varint::VarInt,
};
use std::{
fmt,
net::{Ipv4Addr, SocketAddr},
sync::{
atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering},
Arc, Mutex,
},
time::{Duration, Instant},
};
use zeroize::Zeroizing;
const TLS_EXPORTER_LABEL: &str = "EXPERIMENTAL EXPORTER s2n-quic-dc";
const TLS_EXPORTER_CONTEXT: &str = "";
const TLS_EXPORTER_LENGTH: usize = schedule::EXPORT_SECRET_LEN;
#[derive(Clone)]
pub struct Map {
pub(super) state: Arc<State>,
}
pub(super) struct State {
max_capacity: usize,
rehandshake_period: Duration,
pub(super) peers: flurry::HashMap<SocketAddr, Arc<Entry>>,
pub(super) requested_handshakes: flurry::HashSet<SocketAddr>,
pub(super) ids: flurry::HashMap<Id, Arc<Entry>>,
pub(super) signer: stateless_reset::Signer,
pub(super) control_socket: std::net::UdpSocket,
pub(super) receiver_shared: Arc<receiver::Shared>,
handled_control_packets: AtomicUsize,
cleaner: Cleaner,
}
struct Cleaner {
should_stop: AtomicBool,
thread: Mutex<Option<std::thread::JoinHandle<()>>>,
epoch: AtomicU64,
}
impl Drop for Cleaner {
fn drop(&mut self) {
self.stop();
}
}
impl Cleaner {
fn new() -> Cleaner {
Cleaner {
should_stop: AtomicBool::new(false),
thread: Mutex::new(None),
epoch: AtomicU64::new(1),
}
}
fn stop(&self) {
self.should_stop.store(true, Ordering::Relaxed);
if let Some(thread) =
std::mem::take(&mut *self.thread.lock().unwrap_or_else(|e| e.into_inner()))
{
thread.thread().unpark();
if std::thread::current().id() != thread.thread().id() {
thread.join().unwrap();
}
}
}
fn spawn_thread(&self, state: Arc<State>) {
let state = Arc::downgrade(&state);
let handle = std::thread::spawn(move || loop {
let Some(state) = state.upgrade() else {
break;
};
if state.cleaner.should_stop.load(Ordering::Relaxed) {
break;
}
state.cleaner.clean(&state, EVICTION_CYCLES);
let pause = rand::thread_rng().gen_range(5..60);
drop(state);
std::thread::park_timeout(Duration::from_secs(pause));
});
*self.thread.lock().unwrap() = Some(handle);
}
fn clean(&self, state: &State, eviction_cycles: u64) {
let current_epoch = self.epoch.fetch_add(1, Ordering::Relaxed);
let now = Instant::now();
let mut minimum = u64::MAX;
{
let guard = state.ids.guard();
for (id, entry) in state.ids.iter(&guard) {
let retired_at = entry.retired.0.load(Ordering::Relaxed);
if retired_at == 0 {
minimum = std::cmp::min(entry.used_at.load(Ordering::Relaxed), minimum);
if entry.rehandshake_time() <= now {
state.requested_handshakes.pin().insert(entry.peer);
}
continue;
}
if current_epoch.saturating_sub(retired_at) >= eviction_cycles {
state.ids.remove(id, &guard);
}
}
}
if state.ids.len() <= (state.max_capacity * 95 / 100) {
return;
}
let mut to_remove = std::cmp::max(state.ids.len() / 100, 1);
let guard = state.ids.guard();
for (id, entry) in state.ids.iter(&guard) {
if to_remove > 0 {
if entry.used_at.load(Ordering::Relaxed) == minimum {
state.ids.remove(id, &guard);
to_remove -= 1;
}
} else {
break;
}
}
}
fn epoch(&self) -> u64 {
self.epoch.load(Ordering::Relaxed)
}
}
const EVICTION_CYCLES: u64 = if cfg!(test) { 0 } else { 10 };
impl Map {
pub fn new(signer: stateless_reset::Signer) -> Self {
let control_socket = std::net::UdpSocket::bind((Ipv4Addr::UNSPECIFIED, 0)).unwrap();
control_socket.set_nonblocking(true).unwrap();
let state = State {
max_capacity: 500_000,
rehandshake_period: Duration::from_secs(3600 * 24),
peers: Default::default(),
requested_handshakes: Default::default(),
ids: Default::default(),
cleaner: Cleaner::new(),
signer,
receiver_shared: receiver::Shared::new(),
handled_control_packets: AtomicUsize::new(0),
control_socket,
};
let state = Arc::new(state);
state.cleaner.spawn_thread(state.clone());
Self { state }
}
pub fn secrets_len(&self) -> usize {
self.state.ids.len()
}
pub fn peers_len(&self) -> usize {
self.state.peers.len()
}
pub fn secrets_capacity(&self) -> usize {
self.state.max_capacity
}
pub fn drop_state(&self) {
self.state.peers.pin().clear();
self.state.ids.pin().clear();
}
pub fn contains(&self, peer: SocketAddr) -> bool {
self.state.peers.pin().contains_key(&peer)
&& !self.state.requested_handshakes.pin().contains(&peer)
}
pub fn seal_once(
&self,
peer: SocketAddr,
) -> Option<(seal::Once, Credentials, ApplicationParams)> {
let peers_guard = self.state.peers.guard();
let state = self.state.peers.get(&peer, &peers_guard)?;
state.mark_live(self.state.cleaner.epoch());
let (sealer, credentials) = state.uni_sealer();
Some((sealer, credentials, state.parameters))
}
pub fn open_once(
&self,
credentials: &Credentials,
control_out: &mut Vec<u8>,
) -> Option<open::Once> {
let state = self.pre_authentication(credentials, control_out)?;
let opener = state.uni_opener(self.clone(), credentials);
Some(opener)
}
pub fn pair_for_peer(
&self,
peer: SocketAddr,
features: &TransportFeatures,
) -> Option<(Bidirectional, ApplicationParams)> {
let peers_guard = self.state.peers.guard();
let state = self.state.peers.get(&peer, &peers_guard)?;
state.mark_live(self.state.cleaner.epoch());
let keys = state.bidi_local(features);
Some((keys, state.parameters))
}
pub fn pair_for_credentials(
&self,
credentials: &Credentials,
features: &TransportFeatures,
control_out: &mut Vec<u8>,
) -> Option<(Bidirectional, ApplicationParams)> {
let state = self.pre_authentication(credentials, control_out)?;
let params = state.parameters;
let keys = state.bidi_remote(self.clone(), credentials, features);
Some((keys, params))
}
pub fn handle_unexpected_packet(&self, packet: &Packet) {
match packet {
Packet::Stream(_) => {
}
Packet::Datagram(_) => {
}
Packet::Control(_) => {
}
Packet::StaleKey(packet) => self.handle_control_packet(&(*packet).into()),
Packet::ReplayDetected(packet) => self.handle_control_packet(&(*packet).into()),
Packet::UnknownPathSecret(packet) => self.handle_control_packet(&(*packet).into()),
}
}
pub fn handle_unknown_secret_packet(&self, packet: &control::unknown_path_secret::Packet) {
let ids_guard = self.state.ids.guard();
let Some(state) = self.state.ids.get(packet.credential_id(), &ids_guard) else {
return;
};
if packet.authenticate(&state.sender.stateless_reset).is_none() {
return;
}
self.state
.handled_control_packets
.fetch_add(1, Ordering::Relaxed);
self.state.requested_handshakes.pin().insert(state.peer);
}
pub fn handle_control_packet(&self, packet: &control::Packet) {
if let control::Packet::UnknownPathSecret(ref packet) = &packet {
return self.handle_unknown_secret_packet(packet);
}
let ids_guard = self.state.ids.guard();
let Some(state) = self.state.ids.get(packet.credential_id(), &ids_guard) else {
return;
};
let key = state.sender.control_secret(&state.secret);
match packet {
control::Packet::StaleKey(packet) => {
let Some(packet) = packet.authenticate(key) else {
return;
};
state.mark_live(self.state.cleaner.epoch());
state.sender.update_for_stale_key(packet.min_key_id);
self.state
.handled_control_packets
.fetch_add(1, Ordering::Relaxed);
}
control::Packet::ReplayDetected(packet) => {
let Some(_packet) = packet.authenticate(key) else {
return;
};
self.state
.handled_control_packets
.fetch_add(1, Ordering::Relaxed);
self.state.requested_handshakes.pin().insert(state.peer);
}
control::Packet::UnknownPathSecret(_) => unreachable!(),
}
}
fn pre_authentication(
&self,
identity: &Credentials,
control_out: &mut Vec<u8>,
) -> Option<Arc<Entry>> {
let ids_guard = self.state.ids.guard();
let Some(state) = self.state.ids.get(&identity.id, &ids_guard) else {
let packet = control::UnknownPathSecret {
wire_version: WireVersion::ZERO,
credential_id: identity.id,
};
control_out.resize(control::UnknownPathSecret::PACKET_SIZE, 0);
let stateless_reset = self.state.signer.sign(&identity.id);
let encoder = EncoderBuffer::new(control_out);
packet.encode(encoder, &stateless_reset);
return None;
};
state.mark_live(self.state.cleaner.epoch());
match state.receiver.pre_authentication(identity) {
Ok(()) => {}
Err(e) => {
self.send_control(state, identity, e);
control_out.resize(control::UnknownPathSecret::PACKET_SIZE, 0);
return None;
}
}
Some(state.clone())
}
pub(super) fn insert(&self, entry: Arc<Entry>) {
self.state.requested_handshakes.pin().remove(&entry.peer);
entry.mark_live(self.state.cleaner.epoch());
let id = *entry.secret.id();
let peer = entry.peer;
let ids_guard = self.state.ids.guard();
if self
.state
.ids
.insert(id, entry.clone(), &ids_guard)
.is_some()
{
panic!("inserting a path secret ID twice");
}
let peers_guard = self.state.peers.guard();
if let Some(prev) = self.state.peers.insert(peer, entry, &peers_guard) {
assert_ne!(*prev.secret.id(), id, "duplicate path secret id");
prev.retire(self.state.cleaner.epoch());
}
}
pub(super) fn signer(&self) -> &stateless_reset::Signer {
&self.state.signer
}
#[doc(hidden)]
#[cfg(any(test, feature = "testing"))]
pub fn for_test_with_peers(
peers: Vec<(schedule::Ciphersuite, dc::Version, SocketAddr)>,
) -> (Self, Vec<Id>) {
let provider = Self::new(stateless_reset::Signer::random());
let mut secret = [0; 32];
aws_lc_rs::rand::fill(&mut secret).unwrap();
let mut stateless_reset = [0; control::TAG_LEN];
aws_lc_rs::rand::fill(&mut stateless_reset).unwrap();
let receiver_shared = receiver::Shared::new();
let mut ids = Vec::with_capacity(peers.len());
for (idx, (ciphersuite, version, peer)) in peers.into_iter().enumerate() {
secret[..8].copy_from_slice(&(idx as u64).to_be_bytes()[..]);
stateless_reset[..8].copy_from_slice(&(idx as u64).to_be_bytes()[..]);
let secret = schedule::Secret::new(
ciphersuite,
version,
s2n_quic_core::endpoint::Type::Client,
&secret,
);
ids.push(*secret.id());
let sender = sender::State::new(stateless_reset);
let entry = Entry::new(
peer,
secret,
sender,
receiver_shared.clone().new_receiver(),
dc::testing::TEST_APPLICATION_PARAMS,
dc::testing::TEST_REHANDSHAKE_PERIOD,
);
let entry = Arc::new(entry);
provider.insert(entry);
}
(provider, ids)
}
#[doc(hidden)]
#[cfg(any(test, feature = "testing"))]
pub fn test_insert(&self, peer: SocketAddr) {
let mut secret = [0; 32];
aws_lc_rs::rand::fill(&mut secret).unwrap();
let secret = schedule::Secret::new(
schedule::Ciphersuite::AES_GCM_128_SHA256,
dc::SUPPORTED_VERSIONS[0],
s2n_quic_core::endpoint::Type::Client,
&secret,
);
let sender = sender::State::new([0; control::TAG_LEN]);
let receiver = self.state.receiver_shared.clone().new_receiver();
let entry = Entry::new(
peer,
secret,
sender,
receiver,
dc::testing::TEST_APPLICATION_PARAMS,
dc::testing::TEST_REHANDSHAKE_PERIOD,
);
self.insert(Arc::new(entry));
}
fn send_control(&self, entry: &Entry, credentials: &Credentials, error: receiver::Error) {
let mut buffer = [0; control::MAX_PACKET_SIZE];
let buffer = error.to_packet(entry, credentials, &mut buffer);
let dst = entry.peer;
self.send_control_packet(dst, buffer);
}
pub(crate) fn send_control_packet(&self, dst: SocketAddr, buffer: &[u8]) {
match self.state.control_socket.send_to(buffer, dst) {
Ok(_) => {
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
}
Err(e) => {
tracing::warn!("Failed to send control packet to {:?}: {:?}", dst, e);
}
}
}
#[doc(hidden)]
#[cfg(any(test, feature = "testing"))]
pub fn handled_control_packets(&self) -> usize {
self.state.handled_control_packets.load(Ordering::Relaxed)
}
}
impl receiver::Error {
pub(super) fn to_packet<'buffer>(
self,
entry: &Entry,
credentials: &Credentials,
buffer: &'buffer mut [u8; control::MAX_PACKET_SIZE],
) -> &'buffer [u8] {
debug_assert_eq!(entry.secret.id(), &credentials.id);
let encoder = EncoderBuffer::new(&mut buffer[..]);
let length = match self {
receiver::Error::AlreadyExists => control::ReplayDetected {
wire_version: WireVersion::ZERO,
credential_id: credentials.id,
rejected_key_id: credentials.key_id,
}
.encode(encoder, &entry.secret.control_sealer()),
receiver::Error::Unknown => control::StaleKey {
wire_version: WireVersion::ZERO,
credential_id: credentials.id,
min_key_id: entry.receiver.minimum_unseen_key_id(),
}
.encode(encoder, &entry.secret.control_sealer()),
};
&buffer[..length]
}
}
#[derive(Debug)]
pub(super) struct Entry {
creation_time: Instant,
rehandshake_delta_secs: u32,
peer: SocketAddr,
secret: schedule::Secret,
retired: IsRetired,
used_at: AtomicU64,
sender: sender::State,
receiver: receiver::State,
parameters: ApplicationParams,
}
#[derive(Default)]
struct IsRetired(AtomicU64);
impl fmt::Debug for IsRetired {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("IsRetired").field(&self.retired()).finish()
}
}
impl IsRetired {
fn retired(&self) -> bool {
self.0.load(Ordering::Relaxed) != 0
}
}
impl Entry {
pub fn new(
peer: SocketAddr,
secret: schedule::Secret,
sender: sender::State,
receiver: receiver::State,
mut parameters: ApplicationParams,
rehandshake_time: Duration,
) -> Self {
parameters.max_datagram_size = parameters
.max_datagram_size
.min(crate::stream::MAX_DATAGRAM_SIZE as _);
assert!(rehandshake_time.as_secs() <= u32::MAX as u64);
Self {
creation_time: Instant::now(),
rehandshake_delta_secs: rand::thread_rng().gen_range(
std::cmp::min(rehandshake_time.as_secs(), 360)..rehandshake_time.as_secs(),
) as u32,
peer,
secret,
retired: Default::default(),
used_at: AtomicU64::new(0),
sender,
receiver,
parameters,
}
}
fn retire(&self, at_epoch: u64) {
self.retired.0.store(at_epoch, Ordering::Relaxed);
}
fn mark_live(&self, at_epoch: u64) {
self.used_at.store(at_epoch, Ordering::Relaxed);
}
fn uni_sealer(&self) -> (seal::Once, Credentials) {
let key_id = self.sender.next_key_id();
let credentials = Credentials {
id: *self.secret.id(),
key_id,
};
let sealer = self.secret.application_sealer(key_id);
let sealer = seal::Once::new(sealer);
(sealer, credentials)
}
fn uni_opener(self: Arc<Self>, map: Map, credentials: &Credentials) -> open::Once {
let key_id = credentials.key_id;
let opener = self.secret.application_opener(key_id);
let dedup = Dedup::new(self, key_id, map);
open::Once::new(opener, dedup)
}
fn bidi_local(&self, features: &TransportFeatures) -> Bidirectional {
let key_id = self.sender.next_key_id();
let initiator = Initiator::Local;
let application = ApplicationPair::new(
&self.secret,
key_id,
initiator,
Dedup::disabled(),
);
let control = if features.is_reliable() {
None
} else {
Some(ControlPair::new(&self.secret, key_id, initiator))
};
Bidirectional {
credentials: Credentials {
id: *self.secret.id(),
key_id,
},
application,
control,
}
}
fn bidi_remote(
self: &Arc<Self>,
map: Map,
credentials: &Credentials,
features: &TransportFeatures,
) -> Bidirectional {
let key_id = credentials.key_id;
let initiator = Initiator::Remote;
let application = ApplicationPair::new(
&self.secret,
key_id,
initiator,
Dedup::new(self.clone(), key_id, map),
);
let control = if features.is_reliable() {
None
} else {
Some(ControlPair::new(&self.secret, key_id, initiator))
};
Bidirectional {
credentials: *credentials,
application,
control,
}
}
fn rehandshake_time(&self) -> Instant {
self.creation_time + Duration::from_secs(u64::from(self.rehandshake_delta_secs))
}
}
pub struct Bidirectional {
pub credentials: Credentials,
pub application: ApplicationPair,
pub control: Option<ControlPair>,
}
pub struct ApplicationPair {
pub sealer: seal::Application,
pub opener: open::Application,
}
impl ApplicationPair {
fn new(secret: &schedule::Secret, key_id: VarInt, initiator: Initiator, dedup: Dedup) -> Self {
let (sealer, sealer_ku, opener, opener_ku) = secret.application_pair(key_id, initiator);
let sealer = seal::Application::new(sealer, sealer_ku);
let opener = open::Application::new(opener, opener_ku, dedup);
Self { sealer, opener }
}
}
pub struct ControlPair {
pub sealer: seal::control::Stream,
pub opener: open::control::Stream,
}
impl ControlPair {
fn new(secret: &schedule::Secret, key_id: VarInt, initiator: Initiator) -> Self {
let (sealer, opener) = secret.control_pair(key_id, initiator);
Self { sealer, opener }
}
}
pub struct Dedup {
cell: once_cell::sync::OnceCell<crypto::open::Result>,
init: core::cell::Cell<Option<(Arc<Entry>, VarInt, Map)>>,
}
unsafe impl Sync for Dedup {}
impl Dedup {
#[inline]
fn new(entry: Arc<Entry>, key_id: VarInt, map: Map) -> Self {
Self {
cell: Default::default(),
init: core::cell::Cell::new(Some((entry, key_id, map))),
}
}
#[inline]
pub(crate) fn disabled() -> Self {
Self {
cell: once_cell::sync::OnceCell::with_value(Ok(())),
init: core::cell::Cell::new(None),
}
}
#[inline]
pub(crate) fn disable(&self) {
}
#[inline]
pub fn check(&self) -> crypto::open::Result {
*self.cell.get_or_init(|| {
match self.init.take() {
Some((entry, key_id, map)) => {
let creds = &Credentials {
id: *entry.secret.id(),
key_id,
};
match entry.receiver.post_authentication(creds) {
Ok(()) => Ok(()),
Err(receiver::Error::AlreadyExists) => {
map.send_control(&entry, creds, receiver::Error::AlreadyExists);
Err(crypto::open::Error::ReplayDefinitelyDetected)
}
Err(receiver::Error::Unknown) => {
map.send_control(&entry, creds, receiver::Error::Unknown);
Err(crypto::open::Error::ReplayPotentiallyDetected {
gap: Some(
(*entry.receiver.minimum_unseen_key_id())
.saturating_sub(*creds.key_id),
),
})
}
}
}
None => {
Err(crypto::open::Error::ReplayPotentiallyDetected { gap: None })
}
}
})
}
}
impl fmt::Debug for Dedup {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("Dedup").field("cell", &self.cell).finish()
}
}
pub struct HandshakingPath {
peer: SocketAddr,
dc_version: dc::Version,
parameters: ApplicationParams,
endpoint_type: s2n_quic_core::endpoint::Type,
secret: Option<schedule::Secret>,
map: Map,
}
impl HandshakingPath {
fn new(connection_info: &dc::ConnectionInfo, map: Map) -> Self {
let endpoint_type = match connection_info.endpoint_type {
EndpointType::Server { .. } => s2n_quic_core::endpoint::Type::Server,
EndpointType::Client { .. } => s2n_quic_core::endpoint::Type::Client,
};
Self {
peer: connection_info.remote_address.clone().into(),
dc_version: connection_info.dc_version,
parameters: connection_info.application_params,
endpoint_type,
secret: None,
map,
}
}
}
impl dc::Endpoint for Map {
type Path = HandshakingPath;
fn new_path(&mut self, connection_info: &dc::ConnectionInfo) -> Option<Self::Path> {
Some(HandshakingPath::new(connection_info, self.clone()))
}
fn on_possible_secret_control_packet(
&mut self,
_datagram_info: &DatagramInfo,
payload: &mut [u8],
) -> bool {
let payload = s2n_codec::DecoderBufferMut::new(payload);
match control::Packet::decode(payload) {
Ok((packet, tail)) => {
ensure!(tail.is_empty(), false);
self.handle_control_packet(&packet);
true
}
Err(_) => false,
}
}
}
impl dc::Path for HandshakingPath {
fn on_path_secrets_ready(
&mut self,
session: &impl s2n_quic_core::crypto::tls::TlsSession,
) -> Result<Vec<s2n_quic_core::stateless_reset::Token>, s2n_quic_core::transport::Error> {
let mut material = Zeroizing::new([0; TLS_EXPORTER_LENGTH]);
session
.tls_exporter(
TLS_EXPORTER_LABEL.as_bytes(),
TLS_EXPORTER_CONTEXT.as_bytes(),
&mut *material,
)
.unwrap();
let cipher_suite = match session.cipher_suite() {
s2n_quic_core::crypto::tls::CipherSuite::TLS_AES_128_GCM_SHA256 => {
schedule::Ciphersuite::AES_GCM_128_SHA256
}
s2n_quic_core::crypto::tls::CipherSuite::TLS_AES_256_GCM_SHA384 => {
schedule::Ciphersuite::AES_GCM_256_SHA384
}
_ => return Err(s2n_quic_core::transport::Error::INTERNAL_ERROR),
};
let secret =
schedule::Secret::new(cipher_suite, self.dc_version, self.endpoint_type, &material);
let stateless_reset = self.map.signer().sign(secret.id());
self.secret = Some(secret);
Ok(vec![stateless_reset.into()])
}
fn on_peer_stateless_reset_tokens<'a>(
&mut self,
stateless_reset_tokens: impl Iterator<Item = &'a s2n_quic_core::stateless_reset::Token>,
) {
let sender = sender::State::new(
stateless_reset_tokens
.into_iter()
.next()
.unwrap()
.into_inner(),
);
let receiver = self.map.state.receiver_shared.clone().new_receiver();
let entry = Entry::new(
self.peer,
self.secret
.take()
.expect("peer tokens are only received after secrets are ready"),
sender,
receiver,
self.parameters,
self.map.state.rehandshake_period,
);
let entry = Arc::new(entry);
self.map.insert(entry);
}
}
#[cfg(test)]
mod test;