use super::*;
use crate::{
event::{testing, tracing},
path::secret::{schedule, sender},
};
use s2n_quic_core::{dc, time::NoopClock as Clock};
use std::{
collections::HashSet,
fmt,
net::{Ipv4Addr, SocketAddr, SocketAddrV4},
};
fn fake_entry(port: u16) -> Arc<Entry> {
Entry::fake((Ipv4Addr::LOCALHOST, port).into(), None)
}
#[test]
fn cleans_after_delay() {
let signer = stateless_reset::Signer::new(b"secret");
let map = State::builder()
.with_signer(signer)
.with_capacity(50)
.with_clock(Clock)
.with_subscriber(tracing::Subscriber::default())
.build()
.unwrap();
map.cleaner.stop();
let first = fake_entry(1);
let second = fake_entry(1);
let third = fake_entry(1);
map.test_insert(first.clone());
map.test_insert(second.clone());
assert!(map.ids.contains_key(first.id()));
assert!(map.ids.contains_key(second.id()));
map.cleaner.clean(&map, 1);
map.cleaner.clean(&map, 1);
map.test_insert(third.clone());
assert!(!map.ids.contains_key(first.id()));
assert!(map.ids.contains_key(second.id()));
assert!(map.ids.contains_key(third.id()));
}
#[test]
fn thread_shutdown() {
let signer = stateless_reset::Signer::new(b"secret");
let map = State::builder()
.with_signer(signer)
.with_capacity(10)
.with_clock(Clock)
.with_subscriber((
tracing::Subscriber::default(),
testing::Subscriber::snapshot(),
))
.build()
.unwrap();
let state = Arc::downgrade(&map);
drop(map);
let iterations = 10;
let max_time = core::time::Duration::from_secs(2);
for _ in 0..iterations {
if state.strong_count() == 0 {
return;
}
std::thread::sleep(max_time / iterations);
}
panic!("thread did not shut down after {max_time:?}");
}
#[test]
fn serialize_to_disk_writes_configured_entries() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("secrets");
let signer = stateless_reset::Signer::new(b"secret");
let map = State::builder()
.with_signer(signer)
.with_capacity(50)
.with_clock(Clock)
.with_subscriber(tracing::Subscriber::default())
.with_serializer(disk::Serializer::builder(&path).build().unwrap())
.build()
.unwrap();
map.cleaner.stop();
let first = fake_entry(1);
let second = fake_entry(2);
map.test_insert(first.clone());
map.test_insert(second.clone());
map.serialize_to_disk().unwrap();
let mut decoded: Vec<SocketAddr> = disk::deserialize(&path)
.unwrap()
.map(|e| e.unwrap().peer)
.collect();
decoded.sort();
let mut expected = vec![*first.peer(), *second.peer()];
expected.sort();
assert_eq!(decoded, expected);
}
#[test]
fn serialize_to_disk_emits_event() {
use std::sync::atomic::Ordering;
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("secrets");
let subscriber = Arc::new(testing::Subscriber::no_snapshot());
let signer = stateless_reset::Signer::new(b"secret");
let map = State::builder()
.with_signer(signer)
.with_capacity(50)
.with_clock(Clock)
.with_subscriber(subscriber.clone())
.with_serializer(disk::Serializer::builder(&path).build().unwrap())
.build()
.unwrap();
map.cleaner.stop();
map.test_insert(fake_entry(1));
map.test_insert(fake_entry(2));
map.serialize_to_disk().unwrap();
assert_eq!(
subscriber
.path_secret_map_serialized
.load(Ordering::Relaxed),
1
);
}
#[test]
fn serialize_to_disk_without_serializer_is_noop() {
let signer = stateless_reset::Signer::new(b"secret");
let map = State::builder()
.with_signer(signer)
.with_capacity(50)
.with_clock(Clock)
.with_subscriber(tracing::Subscriber::default())
.build()
.unwrap();
map.cleaner.stop();
map.serialize_to_disk().unwrap();
}
#[derive(Debug, Default)]
struct Model {
invariants: HashSet<Invariant>,
}
#[derive(bolero::TypeGenerator, Debug, Copy, Clone)]
enum Operation {
Insert { ip: u8, path_secret_id: TestId },
AdvanceTime,
ReceiveUnknown { path_secret_id: TestId },
}
#[derive(bolero::TypeGenerator, PartialEq, Eq, Hash, Copy, Clone)]
struct TestId(u8);
impl fmt::Debug for TestId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("TestId")
.field(&self.0)
.field(&self.id())
.finish()
}
}
impl TestId {
fn secret(self) -> schedule::Secret {
let mut export_secret = [0; 32];
export_secret[0] = self.0;
schedule::Secret::new(
schedule::Ciphersuite::AES_GCM_128_SHA256,
dc::SUPPORTED_VERSIONS[0],
s2n_quic_core::endpoint::Type::Client,
&export_secret,
)
}
fn id(self) -> Id {
*self.secret().id()
}
}
#[derive(Debug, PartialEq, Eq, Hash, Copy, Clone)]
enum Invariant {
ContainsIp(SocketAddr),
ContainsId(Id),
IdRemoved(Id),
}
impl Model {
fn perform(&mut self, operation: Operation, state: &State<Clock, tracing::Subscriber>) {
match operation {
Operation::Insert { ip, path_secret_id } => {
let ip = SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::from([0, 0, 0, ip]), 0));
let secret = path_secret_id.secret();
let id = *secret.id();
let stateless_reset = state.signer().sign(&id);
state.test_insert(Arc::new(Entry::new(
ip,
secret,
sender::State::new(stateless_reset),
receiver::State::new(),
dc::testing::TEST_APPLICATION_PARAMS,
dc::testing::TEST_REHANDSHAKE_PERIOD,
None,
)));
self.invariants.insert(Invariant::ContainsIp(ip));
self.invariants.insert(Invariant::ContainsId(id));
}
Operation::AdvanceTime => {
let mut invalidated = Vec::new();
self.invariants.retain(|invariant| {
if let Invariant::ContainsId(id) = invariant {
if state
.get_by_id_untracked(id)
.is_none_or(|v| v.retired_at().is_some())
{
invalidated.push(*id);
return false;
}
}
true
});
for id in invalidated {
assert!(self.invariants.insert(Invariant::IdRemoved(id)), "{id:?}");
}
state.cleaner.clean(state, 0);
}
Operation::ReceiveUnknown { path_secret_id } => {
let id = path_secret_id.id();
let stateless_reset = state.signer.sign(&id);
let packet =
crate::packet::secret_control::unknown_path_secret::Packet::new_for_test(
id,
&stateless_reset,
);
state
.handle_unknown_path_secret_packet(&packet, &"127.0.0.1:1234".parse().unwrap());
if state.should_evict_on_unknown_path_secret()
&& self.invariants.contains(&Invariant::ContainsId(id))
{
self.invariants.retain(|invariant| {
if let Invariant::ContainsId(prev_id) = invariant {
if prev_id == &id {
return false;
}
}
true
});
self.invariants.insert(Invariant::IdRemoved(id));
}
}
}
}
fn check_invariants(&self, state: &State<Clock, tracing::Subscriber>) {
for invariant in self.invariants.iter() {
match invariant {
Invariant::ContainsIp(ip) => {
if state.max_capacity != 5 {
assert!(state.peers.contains_key(ip), "{ip:?}");
}
}
Invariant::ContainsId(id) => {
if state.max_capacity != 5 {
assert!(state.ids.contains_key(id), "{id:?}");
}
}
Invariant::IdRemoved(id) => {
assert!(!state.ids.contains_key(id), "{:?}", state.ids.get(*id));
}
}
}
}
}
fn has_duplicate_pids(ops: &[Operation]) -> bool {
let mut ids = HashSet::new();
for op in ops.iter() {
match op {
Operation::Insert {
ip: _,
path_secret_id,
} => {
if !ids.insert(path_secret_id) {
return true;
}
}
Operation::AdvanceTime => {}
Operation::ReceiveUnknown { path_secret_id: _ } => {
}
}
}
false
}
fn check_invariants_inner(should_evict_on_unknown_path_secret: bool) {
bolero::check!()
.with_type::<Vec<Operation>>()
.with_iterations(10_000)
.for_each(|input: &Vec<Operation>| {
if has_duplicate_pids(input) {
return;
}
let mut model = Model::default();
let signer = stateless_reset::Signer::new(b"secret");
let mut map = State::builder()
.with_signer(signer)
.with_capacity(10_000)
.with_evict_on_unknown_path_secret(should_evict_on_unknown_path_secret)
.with_clock(Clock)
.with_subscriber(tracing::Subscriber::default())
.build()
.unwrap();
map.cleaner.stop();
Arc::<State<Clock, tracing::Subscriber>>::get_mut(&mut map)
.unwrap()
.set_max_capacity(5);
model.check_invariants(&map);
for op in input {
model.perform(*op, &map);
model.check_invariants(&map);
}
})
}
#[test]
fn check_invariants() {
check_invariants_inner(false);
}
#[test]
fn check_invariants_evict_unknown_pid() {
check_invariants_inner(true);
}
#[test]
#[ignore = "fixed size maps currently break overflow assumptions, too small bucket size"]
fn check_invariants_no_overflow() {
bolero::check!()
.with_type::<Vec<Operation>>()
.with_iterations(10_000)
.for_each(|input: &Vec<Operation>| {
if has_duplicate_pids(input) {
return;
}
let mut model = Model::default();
let signer = stateless_reset::Signer::new(b"secret");
let map = State::builder()
.with_signer(signer)
.with_capacity(10_000)
.with_clock(Clock)
.with_subscriber(tracing::Subscriber::default())
.build()
.unwrap();
map.cleaner.stop();
model.check_invariants(&map);
for op in input {
model.perform(*op, &map);
model.check_invariants(&map);
}
})
}
#[test]
#[ignore = "memory growth takes a long time to run"]
fn no_memory_growth() {
let signer = stateless_reset::Signer::new(b"secret");
let map = State::builder()
.with_signer(signer)
.with_capacity(100_000)
.with_clock(Clock)
.with_subscriber(tracing::Subscriber::default())
.build()
.unwrap();
map.cleaner.stop();
for idx in 0..500_000 {
map.test_insert(fake_entry(idx as u16));
}
}
#[test]
fn unknown_path_secret_evicts() {
let signer = stateless_reset::Signer::new(b"secret");
let map = State::builder()
.with_signer(signer)
.with_capacity(5)
.with_evict_on_unknown_path_secret(true)
.with_clock(Clock)
.with_subscriber(tracing::Subscriber::default())
.build()
.unwrap();
let entry = fake_entry(0);
map.test_insert(entry.clone());
let packet = crate::packet::secret_control::unknown_path_secret::Packet::new_for_test(
*entry.clone().id(),
&entry.sender().stateless_reset,
);
assert!(map.ids.contains_key(entry.id()), "{:?}", map.ids);
assert!(map.peers.contains_key(entry.peer()), "{:?}", map.peers);
map.handle_unknown_path_secret_packet(&packet, &"127.0.0.1:1234".parse().unwrap());
assert!(!map.ids.contains_key(entry.id()), "{:?}", map.ids);
assert!(!map.peers.contains_key(entry.peer()), "{:?}", map.peers);
}