use super::*;
pub(super) struct Timer {
callback: Mutex<Option<Box<dyn FnOnce() + Send>>>,
}
impl Timer {
pub(super) fn new(callback: impl FnOnce() + Send + 'static) -> Arc<Self> {
Arc::new(Self {
callback: Mutex::new(Some(Box::new(callback))),
})
}
pub(super) fn fire(&self) {
let callback = self.callback.lock().take();
run(callback);
}
}
impl Drop for Timer {
fn drop(&mut self) {
let callback = self.callback.get_mut().take();
run(callback);
}
}
fn run(callback: Option<Box<dyn FnOnce() + Send>>) {
if let Some(callback) = callback {
callback();
}
}
enum Link {
Deliver(#[allow(dead_code)] Arc<Timer>),
Relay {
level: u8,
slot: u16,
next: Box<Link>,
},
}
type Bucket = Vec<Link>;
pub(super) struct TimingWheel {
ticks: u64,
radix: u64,
levels: Vec<VecDeque<Bucket>>,
}
impl TimingWheel {
pub(super) fn new(max_delay: u64, radix: u64) -> Self {
assert!(radix > 1, "a level needs at least two buckets");
let mut levels = 1_usize;
let mut span = radix;
while span < max_delay {
levels += 1;
span = span.saturating_mul(radix);
}
Self {
ticks: 0,
radix,
levels: (0..levels)
.map(|_| {
std::iter::repeat_with(Bucket::new)
.take(radix as usize)
.collect()
})
.collect(),
}
}
pub(super) fn max_delay(&self) -> u64 {
self.period(self.levels.len())
}
pub(super) fn schedule(&mut self, delay: u64, timer: Arc<Timer>) {
if delay == 0 || delay > self.max_delay() {
return;
}
let deliver = Link::Deliver(timer);
let Some((level, slot, remaining)) = (0..self.levels.len())
.rev()
.find_map(|level| self.entry_leg(level, self.ticks + delay))
else {
return;
};
let route = match remaining {
0 => deliver,
remaining => self.route(remaining, deliver),
};
self.file(level, slot, route);
}
pub(super) fn tick(&mut self) {
self.ticks += 1;
for level in 0..self.levels.len() {
if !self.ticks.is_multiple_of(self.period(level)) {
break;
}
let bucket = self.rotate(level);
for link in bucket {
match link {
Link::Deliver(timer) => drop(timer),
Link::Relay { level, slot, next } => {
self.file(level as usize, slot as usize, *next)
}
}
}
}
}
fn rotate(&mut self, level: usize) -> Bucket {
let queue = &mut self.levels[level];
let bucket = queue.pop_front().unwrap_or_default();
queue.push_back(Bucket::new());
bucket
}
fn file(&mut self, level: usize, slot: usize, link: Link) {
if let Some(bucket) = self.levels[level].get_mut(slot) {
bucket.push(link);
}
}
fn route(&self, remaining: u64, inner: Link) -> Link {
let (level, slot, rest) = self.next_leg(remaining);
let next = match rest {
0 => inner,
rest => self.route(rest, inner),
};
Link::Relay {
level: level as u8,
slot: slot as u16,
next: Box::new(next),
}
}
fn entry_leg(&self, level: usize, target: u64) -> Option<(usize, usize, u64)> {
let period = self.period(level);
let next_rotation = self.ticks - self.ticks % period + period;
let ahead = target.checked_sub(next_rotation)?;
let slot = ahead / period;
(slot < self.radix).then_some((level, slot as usize, ahead % period))
}
fn next_leg(&self, remaining: u64) -> (usize, usize, u64) {
let level = (0..self.levels.len())
.rev()
.find(|level| self.period(*level) <= remaining)
.unwrap_or(0);
let period = self.period(level);
(level, (remaining / period - 1) as usize, remaining % period)
}
fn period(&self, level: usize) -> u64 {
self.radix.saturating_pow(level as u32)
}
}
#[cfg(test)]
mod tests {
use super::*;
struct Harness {
wheel: TimingWheel,
fired: Arc<Mutex<Vec<(u32, u64)>>>,
ticks: Arc<AtomicU64>,
}
impl Harness {
fn new(max_delay: u64, radix: u64) -> Self {
Self {
wheel: TimingWheel::new(max_delay, radix),
fired: Arc::new(Mutex::new(Vec::new())),
ticks: Arc::new(AtomicU64::new(0)),
}
}
fn timer(&self, id: u32) -> Arc<Timer> {
let fired = self.fired.clone();
let ticks = self.ticks.clone();
Timer::new(move || {
fired.lock().push((id, ticks.load(Ordering::Acquire)));
})
}
fn schedule(&mut self, id: u32, delay: u64) {
let timer = self.timer(id);
self.wheel.schedule(delay, timer);
}
fn tick(&mut self) {
self.ticks.fetch_add(1, Ordering::AcqRel);
self.wheel.tick();
}
fn fired_at(&self, id: u32) -> Option<u64> {
self.fired
.lock()
.iter()
.find(|(fired, _)| *fired == id)
.map(|(_, at)| *at)
}
}
#[test]
fn every_delay_fires_on_its_exact_tick() {
let radix = 4;
let max_delay = radix * radix * radix;
for offset in 0..2 * radix * radix {
let mut harness = Harness::new(max_delay, radix);
for _ in 0..offset {
harness.tick();
}
for delay in 1..=max_delay {
harness.schedule(delay as u32, delay);
}
for _ in 0..max_delay {
harness.tick();
}
for delay in 1..=max_delay {
assert_eq!(
harness.fired_at(delay as u32),
Some(offset + delay),
"radix {radix}, offset {offset}, delay {delay}"
);
}
}
}
#[test]
fn every_short_delay_fires_on_its_exact_tick_at_radix_64() {
let radix = 64;
let mut harness = Harness::new(radix * radix * radix * radix, radix);
for _ in 0..100 {
harness.tick();
}
let delays = (1..=200).chain([radix - 1, radix, radix + 1, radix * radix, 4095, 4096]);
for delay in delays.clone() {
harness.schedule(delay as u32, delay);
}
for _ in 0..5000 {
harness.tick();
}
for delay in delays {
assert_eq!(
harness.fired_at(delay as u32),
Some(100 + delay),
"delay {delay}"
);
}
}
#[test]
fn level_count_covers_the_requested_delay() {
assert_eq!(TimingWheel::new(64, 64).max_delay(), 64);
assert_eq!(TimingWheel::new(65, 64).max_delay(), 4096);
assert_eq!(TimingWheel::new(4096, 64).max_delay(), 4096);
assert_eq!(TimingWheel::new(4097, 64).max_delay(), 262_144);
}
#[test]
fn a_delay_the_wheel_cannot_place_fires_at_once() {
let mut harness = Harness::new(64, 64);
harness.schedule(1, 65);
assert_eq!(harness.fired_at(1), Some(0));
harness.schedule(2, 0);
assert_eq!(harness.fired_at(2), Some(0));
}
#[test]
fn rescheduling_the_same_timer_defers_it_to_the_later_route() {
let mut harness = Harness::new(4096, 64);
let timer = harness.timer(1);
harness.wheel.schedule(5, timer.clone());
harness.wheel.schedule(20, timer);
for _ in 0..5 {
harness.tick();
}
assert_eq!(
harness.fired_at(1),
None,
"the earlier route must not fire the timer"
);
for _ in 5..20 {
harness.tick();
}
assert_eq!(harness.fired_at(1), Some(20));
}
#[test]
fn firing_early_makes_the_scheduled_route_inert() {
let mut harness = Harness::new(4096, 64);
let timer = harness.timer(1);
harness.wheel.schedule(10, timer.clone());
timer.fire();
assert_eq!(harness.fired_at(1), Some(0));
for _ in 0..10 {
harness.tick();
}
assert_eq!(harness.fired.lock().len(), 1);
}
#[test]
fn dropping_the_wheel_fires_everything_it_holds() {
let mut harness = Harness::new(262_144, 64);
harness.schedule(1, 5);
harness.schedule(2, 200_000);
let fired = harness.fired.clone();
drop(harness);
let ids = fired
.lock()
.iter()
.map(|(id, _)| *id)
.collect::<HashSet<_>>();
assert_eq!(ids, HashSet::from([1, 2]));
}
}