use std::{
collections::BTreeMap,
time::{Duration, Instant},
};
use futures::future::BoxFuture;
use tokio::sync::mpsc;
use crate::scheduled_timeout::ScheduledTimeoutIdentifier;
macro_rules! some_or_return {
($e: expr) => {
match $e {
Some(v) => v,
None => return,
}
};
}
pub struct SchedulerTask {
scheduled_timeouts: BTreeMap<ScheduledTimeoutIdentifier, BoxFuture<'static, ()>>,
requests_channel_receiver: mpsc::UnboundedReceiver<SchedulerTaskRequest>,
cur_delay: Option<Duration>,
min_sleep_duration: Duration,
}
impl SchedulerTask {
pub fn run(min_sleep_duration: Duration) -> SchedulerTaskHandler {
let (sender, receiver) = mpsc::unbounded_channel();
let mut scheduler_task = SchedulerTask {
scheduled_timeouts: BTreeMap::new(),
requests_channel_receiver: receiver,
cur_delay: None,
min_sleep_duration,
};
tokio::spawn(async move {
scheduler_task.main_loop().await;
});
SchedulerTaskHandler {
schedule_channel_sender: sender,
}
}
async fn main_loop(&mut self) {
loop {
match self.cur_delay {
Some(cur_delay) => {
match tokio::time::timeout(cur_delay, self.requests_channel_receiver.recv())
.await
{
Ok(recv_result) => {
let request = some_or_return!(recv_result);
self.handle_request(request).await;
}
Err(_) => {
let boxed_future =
self.scheduled_timeouts.first_entry().unwrap().remove();
boxed_future.await;
self.update_cur_delay().await;
}
}
}
None => {
let request = some_or_return!(self.requests_channel_receiver.recv().await);
self.handle_request(request).await;
}
}
}
}
async fn handle_request(&mut self, request: SchedulerTaskRequest) {
match request {
SchedulerTaskRequest::ScheduleTimeout {
identifier,
boxed_future,
} => {
self.scheduled_timeouts.insert(identifier, boxed_future);
self.update_cur_delay().await
}
SchedulerTaskRequest::CancelTimeout(cancellation_token) => {
if self
.scheduled_timeouts
.remove(&cancellation_token.timeout_identifier)
.is_some()
{
self.update_cur_delay().await
}
}
}
}
async fn update_cur_delay(&mut self) {
loop {
let delay = match self.scheduled_timeouts.first_key_value() {
Some((timeout_identifier, _boxed_future)) => {
timeout_identifier.get_delay(self.min_sleep_duration)
}
None => {
self.cur_delay = None;
return;
}
};
match delay {
Some(delay) => {
self.cur_delay = Some(delay);
return;
}
None => {
let boxed_future = self.scheduled_timeouts.first_entry().unwrap().remove();
boxed_future.await;
}
}
}
}
}
#[derive(Debug)]
pub struct SchedulerTaskHandler {
schedule_channel_sender: mpsc::UnboundedSender<SchedulerTaskRequest>,
}
impl SchedulerTaskHandler {
pub fn schedule_timeout(
&self,
run_at: Instant,
boxed_future: BoxFuture<'static, ()>,
) -> CancellationToken {
let timeout_identifier = ScheduledTimeoutIdentifier::new(run_at, &boxed_future);
let cancellation_token = CancellationToken {
timeout_identifier: timeout_identifier.clone(),
};
self.send_request(SchedulerTaskRequest::ScheduleTimeout {
identifier: timeout_identifier,
boxed_future,
});
cancellation_token
}
pub fn cancel_timeout(&self, cancellation_token: CancellationToken) {
self.send_request(SchedulerTaskRequest::CancelTimeout(cancellation_token))
}
fn send_request(&self, request: SchedulerTaskRequest) {
if self.schedule_channel_sender.send(request).is_err() {
panic!("scheduler task has stopped unexpectedly")
}
}
}
enum SchedulerTaskRequest {
ScheduleTimeout {
identifier: ScheduledTimeoutIdentifier,
boxed_future: BoxFuture<'static, ()>,
},
CancelTimeout(CancellationToken),
}
#[derive(Debug, Clone)]
pub struct CancellationToken {
timeout_identifier: ScheduledTimeoutIdentifier,
}