use std::cmp::Reverse;
use std::collections::BinaryHeap;
use std::sync::{Arc, Condvar, Mutex, OnceLock};
use std::time::{Duration, Instant};
use super::task_policy::Task;
pub(super) struct TimerService {
state: Arc<TimerState>,
}
struct TimerState {
queue: Mutex<TimerQueue>,
wakeup: Condvar,
}
#[derive(Default)]
struct TimerQueue {
entries: BinaryHeap<Reverse<Entry>>,
}
const MAX_ADVANCE_ROUNDS: usize = 64;
struct Entry {
due: Instant,
seq: u64,
owner: Option<super::RuntimeId>,
task: Task,
}
impl PartialEq for Entry {
fn eq(&self, other: &Self) -> bool {
self.due == other.due && self.seq == other.seq
}
}
impl Eq for Entry {}
impl PartialOrd for Entry {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for Entry {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.due.cmp(&other.due).then(self.seq.cmp(&other.seq))
}
}
impl TimerService {
pub(super) fn global() -> &'static Self {
static TIMER: OnceLock<TimerService> = OnceLock::new();
TIMER.get_or_init(Self::new)
}
fn new() -> Self {
let state = Arc::new(TimerState {
queue: Mutex::new(TimerQueue::default()),
wakeup: Condvar::new(),
});
let worker = Arc::clone(&state);
let _ = std::thread::Builder::new()
.name("tui-lipan-timer".to_string())
.spawn(move || run_timer(&worker));
Self { state }
}
pub(super) fn schedule(&self, delay: Duration, task: Task) {
self.schedule_owned(delay, task, None);
}
pub(super) fn schedule_owned(
&self,
delay: Duration,
task: Task,
owner: Option<super::RuntimeId>,
) {
static SEQ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let seq = SEQ.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let due = Instant::now()
.checked_add(delay)
.unwrap_or_else(Instant::now);
let Ok(mut queue) = self.state.queue.lock() else {
task.cancel();
return;
};
queue.entries.push(Reverse(Entry {
due,
seq,
owner,
task,
}));
drop(queue);
self.state.wakeup.notify_one();
}
pub(super) fn advance_owned(&self, horizon: Instant, owner: super::RuntimeId) -> usize {
let mut ran = 0;
for _ in 0..MAX_ADVANCE_ROUNDS {
let Ok(mut queue) = self.state.queue.lock() else {
break;
};
let mut due = Vec::new();
let mut others = Vec::new();
while matches!(queue.entries.peek(), Some(Reverse(entry)) if entry.due <= horizon) {
let Some(Reverse(entry)) = queue.entries.pop() else {
break;
};
if entry.owner == Some(owner) {
due.push(entry.task);
} else {
others.push(Reverse(entry));
}
}
queue.entries.extend(others);
drop(queue);
if due.is_empty() {
break;
}
for task in due {
task.run();
ran += 1;
}
}
ran
}
}
fn run_timer(state: &Arc<TimerState>) {
loop {
let Ok(mut queue) = state.queue.lock() else {
return;
};
loop {
let now = Instant::now();
let wait = match queue.entries.peek() {
Some(Reverse(entry)) if entry.due <= now => break,
Some(Reverse(entry)) => entry.due.saturating_duration_since(now),
None => Duration::from_secs(3600),
};
let Ok((next, _)) = state.wakeup.wait_timeout(queue, wait) else {
return;
};
queue = next;
}
let mut due = Vec::new();
let now = Instant::now();
while matches!(queue.entries.peek(), Some(Reverse(entry)) if entry.due <= now) {
if let Some(Reverse(entry)) = queue.entries.pop() {
due.push(entry.task);
}
}
drop(queue);
for task in due {
super::TaskExecutor::global().execute(task);
}
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use super::*;
fn service() -> TimerService {
TimerService::new()
}
fn owner(id: u64) -> super::super::RuntimeId {
super::super::RuntimeId::from_raw_for_tests(id)
}
fn counting_task(runs: &Arc<AtomicUsize>) -> Task {
let runs = Arc::clone(runs);
Task::new(move || {
runs.fetch_add(1, Ordering::SeqCst);
})
}
fn horizon(after: Duration) -> Instant {
Instant::now() + after
}
#[test]
fn advance_runs_a_task_once_the_horizon_reaches_it() {
let timer = service();
let runs = Arc::new(AtomicUsize::new(0));
timer.schedule_owned(
Duration::from_millis(200),
counting_task(&runs),
Some(owner(1)),
);
assert_eq!(
timer.advance_owned(horizon(Duration::from_millis(50)), owner(1)),
0
);
assert_eq!(runs.load(Ordering::SeqCst), 0, "not due yet");
assert_eq!(
timer.advance_owned(horizon(Duration::from_millis(250)), owner(1)),
1
);
assert_eq!(runs.load(Ordering::SeqCst), 1);
assert_eq!(
timer.advance_owned(horizon(Duration::from_millis(500)), owner(1)),
0
);
assert_eq!(runs.load(Ordering::SeqCst), 1);
}
#[test]
fn advance_only_claims_timers_the_advancing_runtime_armed() {
let timer = service();
let mine = Arc::new(AtomicUsize::new(0));
let theirs = Arc::new(AtomicUsize::new(0));
timer.schedule_owned(
Duration::from_millis(10),
counting_task(&theirs),
Some(owner(2)),
);
timer.schedule_owned(
Duration::from_millis(20),
counting_task(&mine),
Some(owner(1)),
);
assert_eq!(
timer.advance_owned(horizon(Duration::from_secs(5)), owner(1)),
1
);
assert_eq!(mine.load(Ordering::SeqCst), 1);
assert_eq!(theirs.load(Ordering::SeqCst), 0, "not this runtime's timer");
assert_eq!(
timer.advance_owned(horizon(Duration::from_secs(5)), owner(2)),
1
);
assert_eq!(theirs.load(Ordering::SeqCst), 1);
}
#[test]
fn an_unowned_timer_is_left_to_the_timer_thread() {
let timer = service();
let runs = Arc::new(AtomicUsize::new(0));
timer.schedule(Duration::from_millis(10), counting_task(&runs));
assert_eq!(
timer.advance_owned(horizon(Duration::from_secs(5)), owner(1)),
0
);
assert_eq!(runs.load(Ordering::SeqCst), 0);
}
#[test]
fn advance_resolves_a_chain_within_one_call() {
let timer = service();
let runs = Arc::new(AtomicUsize::new(0));
let state = Arc::clone(&timer.state);
let counter = Arc::clone(&runs);
timer.schedule_owned(
Duration::from_millis(10),
Task::new(move || {
counter.fetch_add(1, Ordering::SeqCst);
let follow_up = counting_task(&counter);
TimerService { state }.schedule_owned(
Duration::from_millis(5),
follow_up,
Some(owner(1)),
);
}),
Some(owner(1)),
);
assert_eq!(
timer.advance_owned(horizon(Duration::from_millis(50)), owner(1)),
2
);
assert_eq!(runs.load(Ordering::SeqCst), 2);
}
#[test]
fn advance_stops_rearming_at_its_round_bound_instead_of_spinning() {
let timer = service();
let runs = Arc::new(AtomicUsize::new(0));
fn arm(state: Arc<TimerState>, runs: Arc<AtomicUsize>) {
let next_state = Arc::clone(&state);
TimerService { state }.schedule_owned(
Duration::from_millis(2),
Task::new(move || {
runs.fetch_add(1, Ordering::SeqCst);
arm(Arc::clone(&next_state), Arc::clone(&runs));
}),
Some(owner(1)),
);
}
arm(Arc::clone(&timer.state), Arc::clone(&runs));
let ran = timer.advance_owned(horizon(Duration::from_millis(50)), owner(1));
assert_eq!(
ran, MAX_ADVANCE_ROUNDS,
"the advance should stop at its bound"
);
assert_eq!(runs.load(Ordering::SeqCst), MAX_ADVANCE_ROUNDS);
}
}