use std::cmp::Reverse;
use std::collections::{BinaryHeap, HashMap};
use std::hash::Hash;
use std::ops::Add;
use std::time::Duration;
use tokio::time::Instant;
#[derive(Debug, PartialEq, Eq)]
struct Entry<K, I> {
deadline: I,
generation: u64,
key: K,
}
impl<K: Eq, I: Ord> Ord for Entry<K, I> {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.deadline.cmp(&other.deadline)
}
}
impl<K: Eq, I: Ord> PartialOrd for Entry<K, I> {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
#[derive(Debug)]
pub struct TimerQueue<K, I = Instant> {
heap: BinaryHeap<Reverse<Entry<K, I>>>,
generations: HashMap<K, u64>,
next_generation: u64,
}
impl<K: Eq, I: Ord> Default for TimerQueue<K, I> {
fn default() -> Self {
Self {
heap: BinaryHeap::new(),
generations: HashMap::new(),
next_generation: 0,
}
}
}
impl<K: Clone + Eq + Hash, I: Ord + Copy + Add<Duration, Output = I>> TimerQueue<K, I> {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn len(&self) -> usize {
self.heap.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.heap.is_empty()
}
pub fn set(&mut self, key: K, now: I, after: Duration) {
let generation = self.bump(&key);
self.heap.push(Reverse(Entry {
deadline: now + after,
generation,
key,
}));
}
pub fn clear(&mut self, key: &K) {
self.bump(key);
}
pub fn forget(&mut self, key: &K) {
self.generations.remove(key);
}
pub fn clear_matching(&mut self, matches: impl Fn(&K) -> bool) {
let keys: Vec<K> = self
.generations
.keys()
.filter(|key| matches(key))
.cloned()
.collect();
for key in keys {
self.bump(&key);
}
}
fn bump(&mut self, key: &K) -> u64 {
if self.next_generation == u64::MAX {
self.compact_generations();
}
self.next_generation += 1;
self.generations.insert(key.clone(), self.next_generation);
self.next_generation
}
fn compact_generations(&mut self) {
let previous = std::mem::take(&mut self.generations);
let entries = std::mem::take(&mut self.heap);
self.next_generation = 0;
for Reverse(mut entry) in entries {
if previous.get(&entry.key) != Some(&entry.generation) {
continue;
}
self.next_generation += 1;
entry.generation = self.next_generation;
self.generations
.insert(entry.key.clone(), self.next_generation);
self.heap.push(Reverse(entry));
}
}
pub fn next_deadline(&mut self) -> Option<I> {
loop {
let Reverse(entry) = self.heap.peek()?;
if self.is_live(entry) {
return Some(entry.deadline);
}
self.heap.pop();
}
}
pub fn take_due(&mut self, now: I) -> Vec<K> {
let mut fired = Vec::new();
while let Some(Reverse(entry)) = self.heap.peek() {
if entry.deadline > now {
break;
}
let Some(Reverse(entry)) = self.heap.pop() else {
break;
};
if !self.is_live(&entry) {
continue;
}
self.bump(&entry.key);
fired.push(entry.key);
}
fired
}
fn is_live(&self, entry: &Entry<K, I>) -> bool {
self.generations
.get(&entry.key)
.is_some_and(|&generation| generation == entry.generation)
}
pub fn forget_matching(&mut self, matches: impl Fn(&K) -> bool) {
self.generations.retain(|key, _| !matches(key));
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::indexing_slicing
)]
mod tests {
use super::*;
use sipx_sip::transaction::{Timer, TransactionKey};
use std::time::Duration;
fn key(branch: &str) -> TransactionKey {
TransactionKey::Rfc3261 {
branch: branch.as_bytes().to_vec(),
sent_by: b"h.example.com".to_vec(),
method: b"INVITE".to_vec(),
}
}
type Transactions = TimerQueue<(TransactionKey, Timer)>;
#[tokio::test(start_paused = true)]
async fn timers_fire_in_deadline_order() {
let mut q = Transactions::new();
let now = Instant::now();
q.set((key("a"), Timer::A), now, Duration::from_millis(500));
q.set((key("b"), Timer::B), now, Duration::from_millis(100));
q.set((key("c"), Timer::E), now, Duration::from_millis(300));
let fired = q.take_due(now + Duration::from_millis(600));
let order: Vec<Timer> = fired.iter().map(|(_, timer)| *timer).collect();
assert_eq!(order, vec![Timer::B, Timer::E, Timer::A]);
}
#[tokio::test]
async fn scheduling_and_firing_need_no_real_time_to_pass() {
let mut q = Transactions::new();
let epoch = Instant::now();
q.set((key("a"), Timer::A), epoch, Duration::from_secs(3600));
assert!(q.take_due(epoch).is_empty(), "not due yet");
assert_eq!(
q.take_due(epoch + Duration::from_secs(3600)).len(),
1,
"an hour later, without an hour passing"
);
}
#[tokio::test(start_paused = true)]
async fn a_cleared_timer_does_not_fire() {
let mut q = Transactions::new();
let now = Instant::now();
q.set((key("a"), Timer::A), now, Duration::from_millis(100));
q.clear(&(key("a"), Timer::A));
assert!(q.take_due(now + Duration::from_millis(200)).is_empty());
}
#[tokio::test(start_paused = true)]
async fn forgetting_one_timer_discards_its_generation_without_scanning_others() {
let mut q = Transactions::new();
let now = Instant::now();
let forgotten = (key("a"), Timer::A);
let live = (key("z"), Timer::B);
q.set(forgotten.clone(), now, Duration::from_millis(100));
q.set(live.clone(), now, Duration::from_millis(200));
q.forget(&forgotten);
assert!(!q.generations.contains_key(&forgotten));
assert!(q.generations.contains_key(&live));
assert_eq!(q.take_due(now + Duration::from_millis(300)), vec![live]);
}
#[tokio::test(start_paused = true)]
async fn reusing_a_forgotten_key_does_not_revive_its_stale_timer() {
let mut q = Transactions::new();
let now = Instant::now();
let reused = (key("a"), Timer::A);
q.set(reused.clone(), now, Duration::from_millis(100));
q.forget(&reused);
q.set(reused.clone(), now, Duration::from_millis(200));
assert!(q.take_due(now + Duration::from_millis(100)).is_empty());
assert_eq!(q.take_due(now + Duration::from_millis(200)), vec![reused]);
}
#[tokio::test(start_paused = true)]
async fn resetting_a_timer_replaces_it() {
let mut q = Transactions::new();
let now = Instant::now();
q.set((key("a"), Timer::A), now, Duration::from_millis(100));
q.set((key("a"), Timer::A), now, Duration::from_millis(500));
assert!(
q.take_due(now + Duration::from_millis(200)).is_empty(),
"the first schedule must not survive"
);
assert_eq!(q.take_due(now + Duration::from_millis(600)).len(), 1);
}
#[tokio::test(start_paused = true)]
async fn clearing_a_transaction_clears_all_of_its_timers() {
let mut q = Transactions::new();
let now = Instant::now();
q.set((key("a"), Timer::A), now, Duration::from_millis(100));
q.set((key("a"), Timer::B), now, Duration::from_millis(200));
q.set((key("z"), Timer::A), now, Duration::from_millis(100));
q.clear_matching(|(k, _)| k == &key("a"));
let fired = q.take_due(now + Duration::from_millis(300));
assert_eq!(fired.len(), 1);
assert_eq!(fired[0].0, key("z"));
}
#[tokio::test(start_paused = true)]
async fn cancelled_entries_do_not_keep_waking_the_loop() {
let mut q = Transactions::new();
let now = Instant::now();
for i in 0..100 {
q.set(
(key(&format!("k{i}")), Timer::A),
now,
Duration::from_millis(10),
);
q.clear(&(key(&format!("k{i}")), Timer::A));
}
q.set((key("live"), Timer::A), now, Duration::from_secs(60));
let deadline = q.next_deadline().expect("a deadline");
assert!(deadline >= now + Duration::from_secs(59));
assert_eq!(q.len(), 1, "stale entries are discarded while looking");
}
#[test]
fn a_virtual_clock_drives_the_queue_with_no_runtime() {
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
struct Virtual(u64);
impl std::ops::Add<Duration> for Virtual {
type Output = Self;
fn add(self, after: Duration) -> Self {
Self(
self.0
.saturating_add(u64::try_from(after.as_millis()).unwrap_or(u64::MAX)),
)
}
}
let mut q: TimerQueue<&'static str, Virtual> = TimerQueue::new();
let epoch = Virtual(0);
q.set("retransmit", epoch, Duration::from_millis(500));
q.set("give-up", epoch, Duration::from_secs(32));
assert!(q.take_due(epoch).is_empty(), "nothing is due at the epoch");
assert_eq!(
q.next_deadline(),
Some(Virtual(500)),
"the queue answers in the caller's own units"
);
assert_eq!(q.take_due(Virtual(500)), vec!["retransmit"]);
assert_eq!(q.take_due(Virtual(31_999)), Vec::<&str>::new());
assert_eq!(q.take_due(Virtual(32_000)), vec!["give-up"]);
}
#[tokio::test(start_paused = true)]
async fn naming_the_queue_without_an_instant_still_means_the_tokio_one() {
let mut q: TimerQueue<(TransactionKey, Timer)> = TimerQueue::new();
let now: Instant = Instant::now();
q.set((key("a"), Timer::A), now, Duration::from_millis(100));
assert_eq!(q.next_deadline(), Some(now + Duration::from_millis(100)));
}
#[tokio::test(start_paused = true)]
async fn the_queue_schedules_any_key_at_all() {
let mut q: TimerQueue<&'static str> = TimerQueue::new();
let now = Instant::now();
q.set("refresh", now, Duration::from_millis(50));
q.set("keepalive", now, Duration::from_millis(10));
assert_eq!(
q.take_due(now + Duration::from_millis(100)),
vec!["keepalive", "refresh"]
);
}
}