use super::{cleaner::Cleaner, entry::ApplicationData, stateless_reset, Entry, Store};
use crate::{
credentials::{Credentials, Id},
crypto,
event::{self, EndpointPublisher as _, IntoEvent as _},
packet::{secret_control as control, Packet},
path::secret::receiver,
};
use s2n_quic_core::{
inet::SocketAddress,
time::{self, Timestamp},
};
use std::{
collections::VecDeque,
hash::BuildHasher,
net::{Ipv4Addr, SocketAddr},
sync::{Arc, Mutex, RwLock, Weak},
time::Duration,
};
#[cfg(test)]
mod tests;
#[derive(Default)]
pub(crate) struct PeerMap(
Mutex<hashbrown::HashTable<Arc<Entry>>>,
std::collections::hash_map::RandomState,
);
#[derive(Default)]
pub(crate) struct IdMap(Mutex<hashbrown::HashTable<Arc<Entry>>>);
impl PeerMap {
fn reserve(&self, additional: usize) {
self.0
.lock()
.unwrap_or_else(|e| e.into_inner())
.reserve(additional, |e| self.hash(e));
}
fn hash(&self, entry: &Entry) -> u64 {
self.hash_key(entry.peer())
}
fn hash_key(&self, entry: &SocketAddr) -> u64 {
self.1.hash_one(entry)
}
pub(crate) fn insert(&self, entry: Arc<Entry>) -> Option<Arc<Entry>> {
let hash = self.hash(&entry);
let mut map = self.0.lock().unwrap_or_else(|e| e.into_inner());
match map.entry(hash, |other| other.peer() == entry.peer(), |e| self.hash(e)) {
hashbrown::hash_table::Entry::Occupied(mut o) => {
Some(std::mem::replace(o.get_mut(), entry))
}
hashbrown::hash_table::Entry::Vacant(v) => {
v.insert(entry);
None
}
}
}
pub(crate) fn contains_key(&self, ip: &SocketAddr) -> bool {
let hash = self.hash_key(ip);
let map = self.0.lock().unwrap_or_else(|e| e.into_inner());
map.find(hash, |o| o.peer() == ip).is_some()
}
pub(crate) fn get(&self, peer: SocketAddr) -> Option<Arc<Entry>> {
let hash = self.hash_key(&peer);
let map = self.0.lock().unwrap_or_else(|e| e.into_inner());
map.find(hash, |o| *o.peer() == peer).cloned()
}
pub(crate) fn clear(&self) {
let mut map = self.0.lock().unwrap_or_else(|e| e.into_inner());
map.clear();
}
pub(super) fn len(&self) -> usize {
let map = self.0.lock().unwrap_or_else(|e| e.into_inner());
map.len()
}
fn remove_exact(&self, entry: &Arc<Entry>) -> Option<Arc<Entry>> {
let hash = self.hash(entry);
let mut map = self.0.lock().unwrap_or_else(|e| e.into_inner());
match map.find_entry(hash, |other| other.id() == entry.id()) {
Ok(o) => Some(o.remove().0),
Err(_) => None,
}
}
}
impl IdMap {
fn reserve(&self, additional: usize) {
self.0
.lock()
.unwrap_or_else(|e| e.into_inner())
.reserve(additional, |e| self.hash(e));
}
fn hash(&self, entry: &Entry) -> u64 {
self.hash_key(entry.id())
}
fn hash_key(&self, entry: &Id) -> u64 {
entry.to_hash()
}
pub(crate) fn insert(&self, entry: Arc<Entry>) -> Option<Arc<Entry>> {
let hash = self.hash(&entry);
let mut map = self.0.lock().unwrap_or_else(|e| e.into_inner());
match map.entry(hash, |other| other.id() == entry.id(), |e| self.hash(e)) {
hashbrown::hash_table::Entry::Occupied(mut o) => {
Some(std::mem::replace(o.get_mut(), entry))
}
hashbrown::hash_table::Entry::Vacant(v) => {
v.insert(entry);
None
}
}
}
#[cfg(test)]
pub(crate) fn contains_key(&self, id: &Id) -> bool {
let hash = self.hash_key(id);
let map = self.0.lock().unwrap_or_else(|e| e.into_inner());
map.find(hash, |o| o.id() == id).is_some()
}
pub(crate) fn get(&self, id: Id) -> Option<Arc<Entry>> {
let hash = self.hash_key(&id);
let map = self.0.lock().unwrap_or_else(|e| e.into_inner());
map.find(hash, |o| *o.id() == id).cloned()
}
pub(crate) fn clear(&self) {
let mut map = self.0.lock().unwrap_or_else(|e| e.into_inner());
map.clear();
}
pub(super) fn len(&self) -> usize {
let map = self.0.lock().unwrap_or_else(|e| e.into_inner());
map.len()
}
pub(super) fn remove(&self, id: Id) -> Option<Arc<Entry>> {
let hash = self.hash_key(&id);
let mut map = self.0.lock().unwrap_or_else(|e| e.into_inner());
match map.find_entry(hash, |other| *other.id() == id) {
Ok(o) => Some(o.remove().0),
Err(_) => None,
}
}
}
pub(super) struct State<C, S>
where
C: 'static + time::Clock + Sync + Send,
S: event::Subscriber,
{
max_capacity: usize,
rehandshake_period: Duration,
pub(super) peers: PeerMap,
pub(super) ids: IdMap,
pub(super) eviction_queue: Mutex<VecDeque<Weak<Entry>>>,
pub(super) signer: stateless_reset::Signer,
pub(super) control_socket: Arc<std::net::UdpSocket>,
#[allow(clippy::type_complexity)]
pub(super) request_handshake: RwLock<Option<Box<dyn Fn(SocketAddr) + Send + Sync>>>,
cleaner: Cleaner,
pub(super) cleaner_peer_seen: PeerMap,
pub(super) rehandshake: Mutex<super::rehandshake::RehandshakeState>,
init_time: Timestamp,
clock: C,
subscriber: S,
#[allow(clippy::type_complexity)]
mk_application_data: RwLock<
Option<
Box<
dyn Fn(&dyn s2n_quic_core::crypto::tls::TlsSession) -> ApplicationData
+ Send
+ Sync,
>,
>,
>,
dummy_application_data: ApplicationData,
}
static CONTROL_SOCKET: Mutex<Weak<std::net::UdpSocket>> = Mutex::new(Weak::new());
impl<C, S> State<C, S>
where
C: 'static + time::Clock + Sync + Send,
S: event::Subscriber,
{
pub fn new(
signer: stateless_reset::Signer,
capacity: usize,
clock: C,
subscriber: S,
) -> Arc<Self> {
let control_socket = {
let mut guard = CONTROL_SOCKET.lock().unwrap();
if let Some(socket) = guard.upgrade() {
socket
} else {
let control_socket = std::net::UdpSocket::bind((Ipv4Addr::UNSPECIFIED, 0)).unwrap();
control_socket.set_nonblocking(true).unwrap();
let control_socket = Arc::new(control_socket);
*guard = Arc::downgrade(&control_socket);
control_socket
}
};
let init_time = clock.get_time();
let rehandshake_period = Duration::from_secs(3600 * 24);
let mut state = Self {
max_capacity: capacity,
rehandshake_period,
peers: Default::default(),
ids: Default::default(),
eviction_queue: Default::default(),
cleaner_peer_seen: Default::default(),
cleaner: Cleaner::new(),
rehandshake: Mutex::new(super::rehandshake::RehandshakeState::new(
rehandshake_period,
)),
signer,
control_socket,
init_time,
clock,
subscriber,
request_handshake: RwLock::new(None),
mk_application_data: RwLock::new(None),
dummy_application_data: Arc::new(()),
};
state.peers.reserve(2 * state.max_capacity);
state.ids.reserve(2 * state.max_capacity);
state.cleaner_peer_seen.reserve(2 * state.max_capacity);
state
.rehandshake
.get_mut()
.unwrap()
.reserve(state.max_capacity);
let state = Arc::new(state);
state.cleaner.spawn_thread(state.clone());
state
.subscriber()
.on_path_secret_map_initialized(event::builder::PathSecretMapInitialized { capacity });
state
}
pub(super) fn evict(&self, evicted: &Arc<Entry>) -> (bool, bool) {
let mut id_removed = false;
let mut peer_removed = false;
if self.ids.remove(*evicted.id()).is_some() {
id_removed = true;
self.subscriber().on_path_secret_map_id_entry_evicted(
event::builder::PathSecretMapIdEntryEvicted {
peer_address: SocketAddress::from(*evicted.peer()).into_event(),
credential_id: evicted.id().into_event(),
age: evicted.age(),
},
);
}
if self.peers.remove_exact(evicted).is_some() {
peer_removed = true;
self.subscriber().on_path_secret_map_address_entry_evicted(
event::builder::PathSecretMapAddressEntryEvicted {
peer_address: SocketAddress::from(*evicted.peer()).into_event(),
credential_id: evicted.id().into_event(),
age: evicted.age(),
},
);
}
(id_removed, peer_removed)
}
pub fn request_handshake(&self, peer: SocketAddr) {
self.subscriber()
.on_path_secret_map_background_handshake_requested(
event::builder::PathSecretMapBackgroundHandshakeRequested {
peer_address: SocketAddress::from(peer).into_event(),
},
);
if let Some(callback) = self
.request_handshake
.read()
.unwrap_or_else(|e| e.into_inner())
.as_deref()
{
(callback)(peer);
}
}
fn register_request_handshake(&self, cb: Box<dyn Fn(SocketAddr) + Send + Sync>) {
*self
.request_handshake
.write()
.unwrap_or_else(|e| e.into_inner()) = Some(cb);
}
fn handle_unknown_secret(
&self,
packet: &control::unknown_path_secret::Packet,
peer: &SocketAddress,
) {
let peer_address = peer.into_event();
self.subscriber().on_unknown_path_secret_packet_received(
event::builder::UnknownPathSecretPacketReceived {
credential_id: packet.credential_id().into_event(),
peer_address,
},
);
let Some(entry) = self.get_by_id_untracked(packet.credential_id()) else {
self.subscriber().on_unknown_path_secret_packet_dropped(
event::builder::UnknownPathSecretPacketDropped {
credential_id: packet.credential_id().into_event(),
peer_address,
},
);
return;
};
if packet
.authenticate(&entry.sender().stateless_reset)
.is_none()
{
self.subscriber().on_unknown_path_secret_packet_rejected(
event::builder::UnknownPathSecretPacketRejected {
credential_id: packet.credential_id().into_event(),
peer_address,
},
);
return;
}
self.subscriber().on_unknown_path_secret_packet_accepted(
event::builder::UnknownPathSecretPacketAccepted {
credential_id: packet.credential_id().into_event(),
peer_address,
},
);
self.request_handshake(*entry.peer());
}
fn handle_stale_key(&self, packet: &control::stale_key::Packet, peer: &SocketAddress) {
let peer_address = peer.into_event();
self.subscriber()
.on_stale_key_packet_received(event::builder::StaleKeyPacketReceived {
credential_id: packet.credential_id().into_event(),
peer_address,
});
let Some(entry) = self.ids.get(*packet.credential_id()) else {
self.subscriber()
.on_stale_key_packet_dropped(event::builder::StaleKeyPacketDropped {
credential_id: packet.credential_id().into_event(),
peer_address,
});
return;
};
let key = entry.control_opener();
let Some(packet) = packet.authenticate(&key) else {
self.subscriber().on_stale_key_packet_rejected(
event::builder::StaleKeyPacketRejected {
credential_id: packet.credential_id().into_event(),
peer_address,
},
);
return;
};
self.subscriber()
.on_stale_key_packet_accepted(event::builder::StaleKeyPacketAccepted {
credential_id: packet.credential_id.into_event(),
peer_address,
});
entry.sender().update_for_stale_key(packet.min_key_id);
}
fn handle_replay_detected(
&self,
packet: &control::replay_detected::Packet,
peer: &SocketAddress,
) {
let peer_address = peer.into_event();
self.subscriber().on_replay_detected_packet_received(
event::builder::ReplayDetectedPacketReceived {
credential_id: packet.credential_id().into_event(),
peer_address,
},
);
let Some(entry) = self.ids.get(*packet.credential_id()) else {
self.subscriber().on_replay_detected_packet_dropped(
event::builder::ReplayDetectedPacketDropped {
credential_id: packet.credential_id().into_event(),
peer_address,
},
);
return;
};
let key = entry.control_opener();
let Some(packet) = packet.authenticate(&key) else {
self.subscriber().on_replay_detected_packet_rejected(
event::builder::ReplayDetectedPacketRejected {
credential_id: packet.credential_id().into_event(),
peer_address,
},
);
return;
};
self.subscriber().on_replay_detected_packet_accepted(
event::builder::ReplayDetectedPacketAccepted {
credential_id: packet.credential_id.into_event(),
key_id: packet.rejected_key_id.into_event(),
peer_address,
},
);
self.request_handshake(*entry.peer());
}
pub fn cleaner(&self) -> &Cleaner {
&self.cleaner
}
#[allow(unused)]
fn set_max_capacity(&mut self, new: usize) {
self.max_capacity = new;
self.peers = Default::default();
self.ids = Default::default();
}
pub(super) fn subscriber(&self) -> event::EndpointPublisherSubscriber<S> {
use event::IntoEvent as _;
let timestamp = self.clock.get_time().into_event();
event::EndpointPublisherSubscriber::new(
event::builder::EndpointMeta { timestamp },
None,
&self.subscriber,
)
}
}
impl<C, S> Store for State<C, S>
where
C: time::Clock + Sync + Send,
S: event::Subscriber,
{
fn secrets_len(&self) -> usize {
self.ids.len()
}
fn peers_len(&self) -> usize {
self.peers.len()
}
fn secrets_capacity(&self) -> usize {
self.max_capacity
}
fn drop_state(&self) {
self.ids.clear();
self.peers.clear();
}
fn contains(&self, peer: &SocketAddr) -> bool {
self.peers.contains_key(peer)
}
fn on_new_path_secrets(&self, entry: Arc<Entry>) {
let id = *entry.id();
let peer = entry.peer();
let same = self.ids.insert(entry.clone());
if same.is_some() {
panic!("inserting a path secret ID twice");
}
{
let mut queue = self
.eviction_queue
.lock()
.unwrap_or_else(|e| e.into_inner());
queue.push_back(Arc::downgrade(&entry));
if queue.len() > self.max_capacity {
let element = queue.pop_front().unwrap();
drop(queue);
if let Some(evicted) = element.upgrade() {
self.evict(&evicted);
}
}
}
self.subscriber().on_path_secret_map_entry_inserted(
event::builder::PathSecretMapEntryInserted {
peer_address: SocketAddress::from(*peer).into_event(),
credential_id: id.into_event(),
},
);
}
fn on_handshake_complete(&self, entry: Arc<Entry>) {
let id = *entry.id();
let peer = *entry.peer();
if let Some(prev) = self.peers.insert(entry.clone()) {
let prev_id = *prev.id();
assert_ne!(prev_id, id, "duplicate path secret id");
prev.retire(self.cleaner.epoch());
self.subscriber().on_path_secret_map_entry_replaced(
event::builder::PathSecretMapEntryReplaced {
peer_address: SocketAddress::from(peer).into_event(),
new_credential_id: id.into_event(),
previous_credential_id: prev_id.into_event(),
},
);
}
self.subscriber()
.on_path_secret_map_entry_ready(event::builder::PathSecretMapEntryReady {
peer_address: SocketAddress::from(peer).into_event(),
credential_id: id.into_event(),
});
}
fn register_request_handshake(&self, cb: Box<dyn Fn(SocketAddr) + Send + Sync>) {
self.register_request_handshake(cb);
}
#[allow(clippy::type_complexity)]
fn register_make_application_data(
&self,
cb: Box<
dyn Fn(&dyn s2n_quic_core::crypto::tls::TlsSession) -> ApplicationData + Send + Sync,
>,
) {
*self
.mk_application_data
.write()
.unwrap_or_else(|e| e.into_inner()) = Some(cb);
}
fn get_by_addr_untracked(&self, peer: &SocketAddr) -> Option<Arc<Entry>> {
self.peers.get(*peer)
}
fn get_by_addr_tracked(&self, peer: &SocketAddr) -> Option<Arc<Entry>> {
let result = self.peers.get(*peer);
self.subscriber().on_path_secret_map_address_cache_accessed(
event::builder::PathSecretMapAddressCacheAccessed {
peer_address: SocketAddress::from(*peer).into_event(),
hit: result.is_some(),
},
);
if let Some(entry) = &result {
entry.set_accessed_addr();
self.subscriber()
.on_path_secret_map_address_cache_accessed_hit(
event::builder::PathSecretMapAddressCacheAccessedHit {
peer_address: SocketAddress::from(*peer).into_event(),
age: entry.age(),
},
);
}
result
}
fn get_by_id_untracked(&self, id: &Id) -> Option<Arc<Entry>> {
self.ids.get(*id)
}
fn get_by_id_tracked(&self, id: &Id) -> Option<Arc<Entry>> {
let result = self.ids.get(*id);
self.subscriber().on_path_secret_map_id_cache_accessed(
event::builder::PathSecretMapIdCacheAccessed {
credential_id: id.into_event(),
hit: result.is_some(),
},
);
if let Some(entry) = &result {
entry.set_accessed_id();
self.subscriber().on_path_secret_map_id_cache_accessed_hit(
event::builder::PathSecretMapIdCacheAccessedHit {
credential_id: id.into_event(),
age: entry.age(),
},
);
}
result
}
fn handle_control_packet(&self, packet: &control::Packet, peer: &SocketAddr) {
match packet {
control::Packet::StaleKey(packet) => self.handle_stale_key(packet, &(*peer).into()),
control::Packet::ReplayDetected(packet) => {
self.handle_replay_detected(packet, &(*peer).into())
}
control::Packet::UnknownPathSecret(packet) => {
self.handle_unknown_secret(packet, &(*peer).into())
}
}
}
fn handle_unexpected_packet(&self, packet: &Packet, peer: &SocketAddr) {
match packet {
Packet::Stream(_) => {
}
Packet::Datagram(_) => {
}
Packet::Control(_) => {
}
Packet::StaleKey(packet) => self.handle_stale_key(packet, &(*peer).into()),
Packet::ReplayDetected(packet) => self.handle_replay_detected(packet, &(*peer).into()),
Packet::UnknownPathSecret(packet) => {
self.handle_unknown_secret(packet, &(*peer).into())
}
}
}
fn signer(&self) -> &stateless_reset::Signer {
&self.signer
}
fn send_control_packet(&self, dst: &SocketAddr, buffer: &mut [u8]) {
match self.control_socket.send_to(buffer, dst) {
Ok(_) => {
match control::Packet::decode(s2n_codec::DecoderBufferMut::new(buffer))
.map(|(t, _)| t)
{
Ok(control::Packet::UnknownPathSecret(packet)) => {
self.subscriber().on_unknown_path_secret_packet_sent(
event::builder::UnknownPathSecretPacketSent {
peer_address: SocketAddress::from(*dst).into_event(),
credential_id: packet.credential_id().into_event(),
},
);
}
Ok(control::Packet::StaleKey(packet)) => {
self.subscriber().on_stale_key_packet_sent(
event::builder::StaleKeyPacketSent {
peer_address: SocketAddress::from(*dst).into_event(),
credential_id: packet.credential_id().into_event(),
},
);
}
Ok(control::Packet::ReplayDetected(packet)) => {
self.subscriber().on_replay_detected_packet_sent(
event::builder::ReplayDetectedPacketSent {
peer_address: SocketAddress::from(*dst).into_event(),
credential_id: packet.credential_id().into_event(),
},
);
}
Err(err) => debug_assert!(false, "decoder error {err:?}"),
}
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
}
Err(e) => {
tracing::warn!("Failed to send control packet to {:?}: {:?}", dst, e);
}
}
}
fn rehandshake_period(&self) -> Duration {
self.rehandshake_period
}
fn check_dedup(
&self,
entry: &Entry,
key_id: s2n_quic_core::varint::VarInt,
) -> crypto::open::Result {
let creds = &Credentials {
id: *entry.id(),
key_id,
};
let starting = *entry.receiver().minimum_unseen_key_id();
match entry.receiver().post_authentication(creds) {
Ok(()) => {
let gap = (*entry.receiver().minimum_unseen_key_id())
.saturating_sub(*creds.key_id);
self.subscriber()
.on_key_accepted(event::builder::KeyAccepted {
credential_id: creds.id.into_event(),
key_id: key_id.into_event(),
gap,
forward_shift: (*creds.key_id).saturating_sub(starting),
});
Ok(())
}
Err(receiver::Error::AlreadyExists) => {
self.send_control_error(entry, creds, receiver::Error::AlreadyExists);
self.subscriber().on_replay_definitely_detected(
event::builder::ReplayDefinitelyDetected {
credential_id: creds.id.into_event(),
key_id: key_id.into_event(),
},
);
Err(crypto::open::Error::ReplayDefinitelyDetected)
}
Err(receiver::Error::Unknown) => {
self.send_control_error(entry, creds, receiver::Error::Unknown);
let gap = (*entry.receiver().minimum_unseen_key_id())
.saturating_sub(*creds.key_id);
self.subscriber().on_replay_potentially_detected(
event::builder::ReplayPotentiallyDetected {
credential_id: creds.id.into_event(),
key_id: key_id.into_event(),
gap,
},
);
Err(crypto::open::Error::ReplayPotentiallyDetected { gap: Some(gap) })
}
}
}
#[cfg(test)]
fn test_stop_cleaner(&self) {
self.cleaner.stop();
}
fn application_data(
&self,
session: &dyn s2n_quic_core::crypto::tls::TlsSession,
) -> ApplicationData {
if let Some(ctxt) = &*self
.mk_application_data
.read()
.unwrap_or_else(|e| e.into_inner())
{
(ctxt)(session)
} else {
self.dummy_application_data.clone()
}
}
}
impl<C, S> Drop for State<C, S>
where
C: 'static + time::Clock + Sync + Send,
S: event::Subscriber,
{
fn drop(&mut self) {
if std::thread::panicking() {
return;
}
let lifetime = self
.clock
.get_time()
.saturating_duration_since(self.init_time);
self.subscriber().on_path_secret_map_uninitialized(
event::builder::PathSecretMapUninitialized {
capacity: self.secrets_capacity(),
entries: self.secrets_len(),
lifetime,
},
);
}
}