use crate::host::error::HostApiError;
use crate::services::timer::timer_manager::TimerManager;
use crate::services::timer::types::{TimerCallback, TimerId, TimerMode};
use async_trait::async_trait;
use dashmap::DashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Notify;
use tokio::time::MissedTickBehavior;
struct TimerEntry {
notify: Arc<Notify>,
}
pub struct TokioTimerManager {
next_id: AtomicU64,
timers: Arc<DashMap<TimerId, TimerEntry>>,
}
impl TokioTimerManager {
pub fn new() -> Self {
Self {
next_id: AtomicU64::new(1),
timers: Arc::new(DashMap::new()),
}
}
}
impl Default for TokioTimerManager {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl TimerManager for TokioTimerManager {
async fn set_timer(
&self,
delay: Duration,
mode: TimerMode,
callback: TimerCallback,
) -> Result<TimerId, HostApiError> {
let id = TimerId::new(self.next_id.fetch_add(1, Ordering::SeqCst));
let notify = Arc::new(Notify::new());
self.timers.insert(
id,
TimerEntry {
notify: notify.clone(),
},
);
let timers = self.timers.clone();
tokio::spawn(async move {
match mode {
TimerMode::OneShot => {
tokio::select! {
_ = tokio::time::sleep(delay) => {
callback(id);
timers.remove(&id);
}
_ = notify.notified() => {
}
}
}
TimerMode::Interval => {
let mut interval = tokio::time::interval(delay);
interval.set_missed_tick_behavior(MissedTickBehavior::Skip);
interval.tick().await;
loop {
tokio::select! {
_ = interval.tick() => {
callback(id);
}
_ = notify.notified() => {
break;
}
}
}
}
}
});
Ok(id)
}
async fn cancel_timer(&self, id: TimerId) -> Result<(), HostApiError> {
if let Some((_, entry)) = self.timers.remove(&id) {
entry.notify.notify_one();
}
Ok(())
}
async fn cancel_all(&self) -> Result<(), HostApiError> {
let entries: Vec<TimerEntry> = self
.timers
.iter()
.map(|e| TimerEntry {
notify: e.notify.clone(),
})
.collect();
self.timers.clear();
for entry in entries {
entry.notify.notify_one();
}
Ok(())
}
}