use std::{
collections::HashMap,
sync::{Arc, Mutex, PoisonError},
time::Duration,
};
use lru_slab::LruSlab;
use crate::VarInt;
const DEFAULT_MAX_SERVER_ENDPOINTS: u32 = 500;
const MIN_INITIAL_RTT: Duration = Duration::from_millis(10);
const MAX_INITIAL_RTT: Duration = Duration::from_secs(1);
pub trait ServerRttStore: Send + Sync {
fn insert(&self, server_name: &str, server_port: u16, smoothed_rtt: Duration);
fn get(&self, server_name: &str, server_port: u16) -> Option<Duration>;
fn remove_if_eq(&self, server_name: &str, server_port: u16, expected_rtt: Duration);
}
#[derive(Debug)]
pub(crate) struct ServerRttMemoryStore(Mutex<State>);
impl ServerRttMemoryStore {
fn new(max_server_endpoints: u32) -> Self {
Self(Mutex::new(State {
max_server_endpoints,
lookup: HashMap::new(),
lru: LruSlab::default(),
}))
}
}
impl ServerRttStore for ServerRttMemoryStore {
#[inline]
fn insert(&self, server_name: &str, server_port: u16, smoothed_rtt: Duration) {
self.0
.lock()
.unwrap_or_else(PoisonError::into_inner)
.insert(server_name, server_port, smoothed_rtt);
}
#[inline]
fn get(&self, server_name: &str, server_port: u16) -> Option<Duration> {
self.0
.lock()
.unwrap_or_else(PoisonError::into_inner)
.get(server_name, server_port)
}
#[inline]
fn remove_if_eq(&self, server_name: &str, server_port: u16, expected_rtt: Duration) {
self.0
.lock()
.unwrap_or_else(PoisonError::into_inner)
.remove_if_eq(server_name, server_port, expected_rtt);
}
}
impl Default for ServerRttMemoryStore {
fn default() -> Self {
Self::new(DEFAULT_MAX_SERVER_ENDPOINTS)
}
}
#[derive(Debug)]
struct State {
max_server_endpoints: u32,
lookup: HashMap<Arc<str>, HashMap<u16, u32>>,
lru: LruSlab<CacheEntry>,
}
impl State {
fn insert(&mut self, server_name: &str, server_port: u16, rtt: Duration) {
if self.max_server_endpoints == 0 {
return;
}
let server_name = match self.lookup.get_key_value(server_name) {
Some((stored_server_name, ports)) => {
if let Some(slab_key) = ports.get(&server_port).copied() {
self.lru.get_mut(slab_key).rtt = rtt;
return;
}
Arc::clone(stored_server_name)
}
None => Arc::<str>::from(server_name),
};
if self.lru.len() >= self.max_server_endpoints {
let Some(slab_key) = self.lru.lru() else {
return;
};
let evicted = self.lru.remove(slab_key);
self.remove_lookup(&evicted.server_name, evicted.server_port);
}
let slab_key = self.lru.insert(CacheEntry {
server_name: Arc::clone(&server_name),
server_port,
rtt,
});
self.lookup
.entry(server_name)
.or_default()
.insert(server_port, slab_key);
}
fn get(&mut self, server_name: &str, server_port: u16) -> Option<Duration> {
let slab_key = self.lookup.get(server_name)?.get(&server_port).copied()?;
Some(self.lru.get_mut(slab_key).rtt)
}
fn remove_if_eq(&mut self, server_name: &str, server_port: u16, expected_rtt: Duration) {
let Some(slab_key) = self
.lookup
.get(server_name)
.and_then(|ports| ports.get(&server_port))
.copied()
else {
return;
};
if self.lru.peek(slab_key).rtt != expected_rtt {
return;
}
if let Some(slab_key) = self.remove_lookup(server_name, server_port) {
self.lru.remove(slab_key);
}
}
fn remove_lookup(&mut self, server_name: &str, server_port: u16) -> Option<u32> {
let (slab_key, remove_server_name) = {
let ports = self.lookup.get_mut(server_name)?;
let slab_key = ports.remove(&server_port)?;
(slab_key, ports.is_empty())
};
if remove_server_name {
self.lookup.remove(server_name);
}
Some(slab_key)
}
}
#[derive(Debug)]
struct CacheEntry {
server_name: Arc<str>,
server_port: u16,
rtt: Duration,
}
pub(crate) fn encode(rtt: Duration) -> Option<(Duration, VarInt)> {
let rtt = sanitize(rtt)?;
let micros = rtt.as_micros() as u32;
Some((rtt, VarInt::from_u32(micros)))
}
pub(crate) fn decode(value: VarInt) -> Option<Duration> {
sanitize(Duration::from_micros(value.into_inner()))
}
fn sanitize(rtt: Duration) -> Option<Duration> {
(!rtt.is_zero()).then(|| rtt.clamp(MIN_INITIAL_RTT, MAX_INITIAL_RTT))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn initial_rtt_is_clamped_to_bounds() {
assert_eq!(encode(Duration::ZERO), None);
assert_eq!(decode(VarInt::from_u32(0)), None);
let (rtt, encoded) = encode(Duration::from_millis(1)).unwrap();
assert_eq!(rtt, MIN_INITIAL_RTT);
assert_eq!(encoded.into_inner(), 10_000);
assert_eq!(decode(VarInt::from_u32(1)), Some(MIN_INITIAL_RTT));
let (rtt, encoded) = encode(Duration::from_millis(20)).unwrap();
assert_eq!(rtt, Duration::from_millis(20));
assert_eq!(decode(encoded), Some(rtt));
let (rtt, encoded) = encode(Duration::from_secs(2)).unwrap();
assert_eq!(rtt, MAX_INITIAL_RTT);
assert_eq!(encoded.into_inner(), 1_000_000);
assert_eq!(decode(VarInt::from_u32(2_000_000)), Some(MAX_INITIAL_RTT));
}
#[test]
fn memory_store_keeps_latest_srtt_per_server_endpoint() {
let store = ServerRttMemoryStore::default();
assert_eq!(store.get("example.com", 443), None);
store.insert("example.com", 443, Duration::from_millis(20));
store.insert("example.com", 443, Duration::from_millis(30));
assert_eq!(
store.get("example.com", 443),
Some(Duration::from_millis(30))
);
assert_eq!(store.get("example.com", 8443), None);
}
#[test]
fn memory_store_evicts_least_recently_used_endpoint() {
let store = ServerRttMemoryStore::new(2);
let rtt = Duration::from_millis(20);
store.insert("first.example", 443, rtt);
store.insert("second.example", 443, rtt);
assert_eq!(store.get("first.example", 443), Some(rtt));
store.insert("third.example", 443, rtt);
assert_eq!(store.get("first.example", 443), Some(rtt));
assert_eq!(store.get("second.example", 443), None);
assert_eq!(store.get("third.example", 443), Some(rtt));
}
#[test]
fn zero_capacity_memory_store_stays_empty() {
let store = ServerRttMemoryStore::new(0);
store.insert("example.com", 443, Duration::from_millis(20));
assert_eq!(store.get("example.com", 443), None);
}
#[test]
fn memory_store_removes_only_the_matching_endpoint_value() {
let store = ServerRttMemoryStore::default();
let rtt = Duration::from_millis(20);
store.insert("example.com", 443, rtt);
store.insert("example.com", 8443, rtt);
store.remove_if_eq("example.com", 443, Duration::from_millis(30));
assert_eq!(store.get("example.com", 443), Some(rtt));
store.remove_if_eq("example.com", 443, rtt);
store.remove_if_eq("example.com", 443, rtt);
assert_eq!(store.get("example.com", 443), None);
assert_eq!(store.get("example.com", 8443), Some(rtt));
}
#[test]
fn memory_store_preserves_a_different_newer_value_during_stale_removal() {
let store = ServerRttMemoryStore::default();
let old_rtt = Duration::from_millis(20);
let new_rtt = Duration::from_millis(30);
store.insert("example.com", 443, old_rtt);
let cached_rtt = store.get("example.com", 443).unwrap();
store.insert("example.com", 443, new_rtt);
store.remove_if_eq("example.com", 443, cached_rtt);
assert_eq!(store.get("example.com", 443), Some(new_rtt));
}
#[test]
fn memory_store_reuses_server_name_allocation_across_ports() {
let store = ServerRttMemoryStore::default();
let rtt = Duration::from_millis(20);
store.insert("example.com", 443, rtt);
store.insert("example.com", 8443, rtt);
let mut state = store.0.lock().unwrap();
let (lookup_name, first_key, second_key) = {
let (server_name, ports) = state.lookup.get_key_value("example.com").unwrap();
(Arc::clone(server_name), ports[&443], ports[&8443])
};
let first_name = Arc::clone(&state.lru.get_mut(first_key).server_name);
let second_name = Arc::clone(&state.lru.get_mut(second_key).server_name);
assert!(Arc::ptr_eq(&lookup_name, &first_name));
assert!(Arc::ptr_eq(&lookup_name, &second_name));
}
}