use super::{cleaner::Cleaner, stateless_reset, Entry, Store};
use crate::{
credentials::{Credentials, Id},
crypto,
event::{self, EndpointPublisher as _, IntoEvent as _},
fixed_map::{self, ReadGuard},
packet::{secret_control as control, Packet},
path::secret::receiver,
};
use s2n_quic_core::{
inet::SocketAddress,
time::{self, Timestamp},
};
use std::{
hash::{BuildHasherDefault, Hasher},
net::{Ipv4Addr, SocketAddr},
sync::{Arc, Mutex, Weak},
time::Duration,
};
#[cfg(test)]
mod tests;
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: fixed_map::Map<SocketAddr, Arc<Entry>>,
pub(super) requested_handshakes: flurry::HashSet<SocketAddr>,
pub(super) ids: fixed_map::Map<Id, Arc<Entry>, BuildHasherDefault<NoopIdHasher>>,
pub(super) signer: stateless_reset::Signer,
pub(super) control_socket: Arc<std::net::UdpSocket>,
pub(super) receiver_shared: Arc<receiver::Shared>,
cleaner: Cleaner,
init_time: Timestamp,
clock: C,
subscriber: S,
}
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 state = Self {
max_capacity: capacity,
rehandshake_period: Duration::from_secs(3600 * 24),
peers: fixed_map::Map::with_capacity(capacity, Default::default()),
ids: fixed_map::Map::with_capacity(capacity, Default::default()),
requested_handshakes: Default::default(),
cleaner: Cleaner::new(),
signer,
receiver_shared: receiver::Shared::new(),
control_socket,
init_time,
clock,
subscriber,
};
let state = Arc::new(state);
state.cleaner.spawn_thread(state.clone());
state
.subscriber()
.on_path_secret_map_initialized(event::builder::PathSecretMapInitialized { capacity });
state
}
pub fn request_handshake(&self, peer: SocketAddr) {
let handshakes = self.requested_handshakes.pin();
if handshakes.len() <= 6000 {
handshakes.insert(peer);
self.subscriber()
.on_path_secret_map_background_handshake_requested(
event::builder::PathSecretMapBackgroundHandshakeRequested {
peer_address: SocketAddress::from(peer).into_event(),
},
);
}
}
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_by_key(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_by_key(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 = fixed_map::Map::with_capacity(new, Default::default());
self.ids = fixed_map::Map::with_capacity(new, 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.count()
}
fn peers_len(&self) -> usize {
self.peers.count()
}
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 needs_handshake(&self, peer: &SocketAddr) -> bool {
self.requested_handshakes.pin().contains(peer)
}
fn on_new_path_secrets(&self, entry: Arc<Entry>) {
let id = *entry.id();
let peer = entry.peer();
self.requested_handshakes.pin().remove(peer);
let (same, other) = self.ids.insert(id, entry.clone());
if same.is_some() {
panic!("inserting a path secret ID twice");
}
if let Some(evicted) = other {
self.subscriber().on_path_secret_map_id_entry_evicted(
event::builder::PathSecretMapIdEntryEvicted {
peer_address: SocketAddress::from(*evicted.1.peer()).into_event(),
credential_id: evicted.1.id().into_event(),
age: evicted.1.age(),
},
);
}
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();
let (same, other) = self.peers.insert(peer, entry);
if let Some(prev) = same {
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(),
},
);
}
if let Some(evicted) = other {
self.subscriber().on_path_secret_map_address_entry_evicted(
event::builder::PathSecretMapAddressEntryEvicted {
peer_address: SocketAddress::from(*evicted.1.peer()).into_event(),
credential_id: evicted.1.id().into_event(),
age: evicted.1.age(),
},
);
}
self.subscriber()
.on_path_secret_map_entry_ready(event::builder::PathSecretMapEntryReady {
peer_address: SocketAddress::from(peer).into_event(),
credential_id: id.into_event(),
});
}
fn get_by_addr_untracked(&self, peer: &SocketAddr) -> Option<ReadGuard<Arc<Entry>>> {
self.peers.get_by_key(peer)
}
fn get_by_addr_tracked(&self, peer: &SocketAddr) -> Option<ReadGuard<Arc<Entry>>> {
let result = self.peers.get_by_key(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<ReadGuard<Arc<Entry>>> {
self.ids.get_by_key(id)
}
fn get_by_id_tracked(&self, id: &Id) -> Option<Arc<Entry>> {
let result = self.ids.get_by_key(id).map(|v| Arc::clone(&v));
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(),
},
);
}
if let Some(entry) = &result {
if let Some(evicted) = self.peers.insert_new_key(*entry.peer(), entry.clone()) {
self.subscriber().on_path_secret_map_address_entry_evicted(
event::builder::PathSecretMapAddressEntryEvicted {
peer_address: SocketAddress::from(*evicted.1.peer()).into_event(),
credential_id: evicted.1.id().into_event(),
age: evicted.1.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 receiver(&self) -> &Arc<receiver::Shared> {
&self.receiver_shared
}
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();
}
}
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,
},
);
}
}
#[derive(Default)]
pub(super) struct NoopIdHasher(Option<u64>);
impl Hasher for NoopIdHasher {
fn finish(&self) -> u64 {
self.0.unwrap()
}
fn write(&mut self, _bytes: &[u8]) {
unimplemented!()
}
fn write_u64(&mut self, x: u64) {
debug_assert!(self.0.is_none());
self.0 = Some(x);
}
}