use crate::errors::{TaskError, TaskResult};
use crate::generator::TaskGenerator;
use crate::task::{run_task, Status, Task, TaskCmd, TaskResponse};
use crate::{scheduler_log, task_log};
use chrono::prelude::*;
use chrono::Utc;
use futures::future::join_all;
use futures::StreamExt;
use std::future::Future;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::sync::{mpsc, oneshot, Notify};
use tokio::task::JoinHandle;
#[derive(Debug, PartialEq)]
pub(crate) enum ExecutionStatus {
Success(usize),
NoExecution,
HadError(usize, usize),
}
#[derive(Debug)]
pub struct TaskHandle {
id: usize,
handle: JoinHandle<()>,
sender: mpsc::Sender<TaskCmd>,
is_init: bool,
}
#[derive(Clone, Debug)]
pub struct SchedulerHandle {
notify: Arc<Notify>,
}
impl SchedulerHandle {
pub fn shutdown(&self) {
self.notify.notify_one();
}
}
pub struct TaskScheduler<T>
where
T: TimeZone + Clone + Send + 'static,
{
handles: Vec<TaskHandle>,
task_gen: Option<TaskGenerator<T>>,
sleep: usize,
next_id: usize,
timezone: T,
shutdown: Arc<Notify>,
}
impl<T> TaskScheduler<T>
where
T: TimeZone + Clone + Send + 'static,
<T as TimeZone>::Offset: Send,
{
pub fn default(timezone: T) -> TaskScheduler<T> {
TaskScheduler {
handles: Vec::new(),
task_gen: None,
sleep: 1000,
timezone,
next_id: 0,
shutdown: Arc::new(Notify::new()),
}
}
pub fn handle(&self) -> SchedulerHandle {
SchedulerHandle {
notify: self.shutdown.clone(),
}
}
pub fn new(sleep: usize, timezone: T) -> TaskScheduler<T> {
TaskScheduler {
sleep,
..TaskScheduler::default(timezone)
}
}
pub fn set_task_gen(&mut self, task_gen: TaskGenerator<T>) -> &mut TaskScheduler<T> {
self.task_gen = Some(task_gen);
self
}
pub fn add_task(
&mut self,
task: TaskResult<Task<T>>,
) -> Result<&mut TaskScheduler<T>, TaskError> {
match task {
Ok(mut task) => {
let (sender, receiver) = mpsc::channel(32);
task.set_receiver(receiver);
task.set_id(self.next_id);
let handle = tokio::spawn(run_task(task));
self.handles.push(TaskHandle {
id: self.next_id,
handle,
sender,
is_init: false,
});
self.next_id += 1;
Ok(self)
}
Err(e) => Err(e),
}
}
pub(crate) async fn execute_tasks(&mut self) -> ExecutionStatus {
let mut receivers: Vec<oneshot::Receiver<TaskResponse>> = Vec::new();
for handle in &self.handles {
let (sender, recv) = oneshot::channel();
let _ = handle.sender.send(TaskCmd::Run { sender }).await;
receivers.push(recv);
}
let err_no: Arc<Mutex<usize>> = Arc::new(Mutex::new(0usize));
let total_runs: Arc<Mutex<usize>> = Arc::new(Mutex::new(0usize));
futures::stream::iter(receivers)
.for_each(|r| async {
let status = match r.await {
Ok(response) => response.status,
Err(_) => {
scheduler_log!(
log::Level::Error,
"A task failed to report its run status and will be skipped"
);
return;
}
};
match status {
Status::Executed => {
*total_runs.lock().unwrap() += 1;
}
Status::Failed => {
*err_no.lock().unwrap() += 1;
*total_runs.lock().unwrap() += 1;
}
_ => { }
};
})
.await;
receivers = Vec::new();
for handle in &self.handles {
let (send, recv) = oneshot::channel();
let _ = handle
.sender
.send(TaskCmd::Reschedule { sender: send })
.await;
receivers.push(recv);
}
for recv in receivers {
let res = match recv.await {
Ok(res) => res,
Err(_) => {
scheduler_log!(
log::Level::Error,
"A task failed to report its reschedule status and will be skipped"
);
continue;
}
};
if res.status == Status::Finished || res.status == Status::ForceRemoved {
for handle in &self.handles {
if handle.id == res.id {
task_log!(
res.id,
log::Level::Debug,
"Removing task due to {}",
if res.status == Status::Finished {
"end of execution cycle"
} else {
"force removal"
}
);
handle.handle.abort();
}
}
let index = self.handles.iter().position(|x| x.id == res.id).unwrap();
self.handles.remove(index);
}
}
if *total_runs.lock().unwrap() > 0 {
if *err_no.lock().unwrap() == 0 {
ExecutionStatus::Success(*total_runs.lock().unwrap())
} else {
ExecutionStatus::HadError(*total_runs.lock().unwrap(), *err_no.lock().unwrap())
}
} else {
ExecutionStatus::NoExecution
}
}
pub(crate) async fn init_tasks(&mut self) {
let mut receivers: Vec<oneshot::Receiver<TaskResponse>> = Vec::new();
let mut count: usize = 0;
for handle in &self.handles {
if !handle.is_init {
let (sender, recv) = oneshot::channel();
let _ = handle.sender.send(TaskCmd::Init { sender }).await;
receivers.push(recv);
count += 1;
}
}
if count > 0 {
join_all(receivers).await.iter().for_each(|r| match r {
Ok(r) => match r.status {
Status::Scheduled => {
self.handles
.iter_mut()
.filter(|h| h.id == r.id)
.for_each(|h| {
task_log!(h.id, log::Level::Info, "Initialized");
h.is_init = true;
});
}
_ => {
task_log!(r.id, log::Level::Error, "Failed to initialize");
}
},
Err(_) => {
scheduler_log!(
log::Level::Error,
"A task failed to report its init status and will be skipped"
);
}
});
}
}
fn run_task_gen(&mut self) -> bool {
match self.task_gen {
Some(ref mut tg) => {
if tg.next_exec <= Utc::now().with_timezone(&self.timezone) {
return match tg.run() {
Some(t) => {
let _ = self.add_task(t);
true
}
None => false,
};
}
false
}
None => false,
}
}
fn shutdown_tasks(&mut self) {
for handle in &self.handles {
handle.handle.abort();
}
self.handles.clear();
}
async fn tick(&mut self) {
if self.run_task_gen() {
self.init_tasks().await;
}
match self.execute_tasks().await {
ExecutionStatus::Success(c) => {
scheduler_log!(
log::Level::Info,
"Execution round completed successfully for {} task(s)",
c
);
}
ExecutionStatus::HadError(c, e) => {
scheduler_log!(
log::Level::Error,
"Execution round ran {} task(s) with {} error(s)",
c,
e
);
}
_ => { }
}
}
pub async fn run(&mut self) {
let shutdown = self.shutdown.clone();
self.run_until(async move { shutdown.notified().await })
.await;
}
pub async fn run_until<F>(&mut self, shutdown: F)
where
F: Future<Output = ()>,
{
scheduler_log!(
log::Level::Info,
"Scheduler started with {} task(s) in queue",
self.handles.len()
);
tokio::pin!(shutdown);
self.init_tasks().await;
let mut interval = tokio::time::interval(Duration::from_millis(self.sleep as u64));
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
biased;
_ = &mut shutdown => {
scheduler_log!(log::Level::Info, "Shutdown requested, stopping scheduler");
break;
}
_ = interval.tick() => {
self.tick().await;
}
}
}
self.shutdown_tasks();
scheduler_log!(log::Level::Info, "Scheduler stopped");
}
}
#[cfg(test)]
mod test {
use super::*;
use crate::task::TaskStepStatusErr::{Error, ErrorDelete};
use crate::task::TaskStepStatusOk::Success;
use crate::TaskBuilder;
use chrono::Local;
use std::time::Duration;
#[tokio::test]
async fn test_scheduler_normal_flow() {
let mut scheduler = TaskScheduler::new(500, Local);
scheduler
.add_task(Task::new("* * * * * * *", None, Some(2), Local))
.unwrap()
.add_task(Task::new("* * * * * * *", None, None, Local))
.unwrap();
assert_eq!(scheduler.handles.len(), 2);
scheduler.init_tasks().await;
tokio::time::sleep(Duration::from_millis(1000)).await;
let status: ExecutionStatus = scheduler.execute_tasks().await;
assert_eq!(status, ExecutionStatus::Success(2));
assert_eq!(scheduler.handles.len(), 2);
tokio::time::sleep(Duration::from_millis(1000)).await;
let status: ExecutionStatus = scheduler.execute_tasks().await;
assert_eq!(status, ExecutionStatus::Success(2));
assert_eq!(scheduler.handles.len(), 1);
}
#[tokio::test]
async fn test_scheduler_normal_force_deletion() {
let mut scheduler = TaskScheduler::new(500, Local);
let mut task = Task::new("* * * * * * *", None, Some(1), Local).unwrap();
task.add_step_default(|| async { Err(ErrorDelete) });
scheduler
.add_task(Ok(task))
.unwrap()
.add_task(Task::new("* * * * * * *", None, None, Local))
.unwrap();
assert_eq!(scheduler.handles.len(), 2);
scheduler.init_tasks().await;
tokio::time::sleep(Duration::from_millis(1000)).await;
scheduler.execute_tasks().await;
assert_eq!(scheduler.handles.len(), 1);
tokio::time::sleep(Duration::from_millis(1000)).await;
scheduler.execute_tasks().await;
assert_eq!(scheduler.handles.len(), 1);
}
#[tokio::test]
async fn test_scheduler_normal_flow_no_execution() {
let mut scheduler = TaskScheduler::new(500, Local);
scheduler.init_tasks().await;
tokio::time::sleep(Duration::from_millis(1000)).await;
let status: ExecutionStatus = scheduler.execute_tasks().await;
assert_eq!(status, ExecutionStatus::NoExecution);
}
#[tokio::test]
async fn test_scheduler_normal_flow_error_case() {
let mut scheduler = TaskScheduler::new(500, Local);
let mut task = Task::new("* * * * * * *", None, Some(1), Local).unwrap();
task.add_step_default(|| async { Ok(Success) });
task.add_step_default(|| async { Err(Error) });
scheduler.add_task(Ok(task)).unwrap();
assert_eq!(scheduler.handles.len(), 1);
scheduler.init_tasks().await;
tokio::time::sleep(Duration::from_millis(1000)).await;
scheduler.execute_tasks().await;
assert_eq!(scheduler.handles.len(), 0);
}
#[tokio::test]
async fn test_scheduler_with_generator() {
let mut scheduler = TaskScheduler::new(500, Local);
scheduler.set_task_gen(TaskGenerator::new("* * * * * * *", Local, || None).unwrap());
assert_eq!(scheduler.handles.len(), 0);
tokio::time::sleep(Duration::from_millis(1000)).await;
scheduler.run_task_gen();
assert_eq!(scheduler.handles.len(), 0);
scheduler.set_task_gen(
TaskGenerator::new("* * * * * * *", Local, || {
Some(TaskBuilder::new(Local).every("* * * * * * *").build())
})
.unwrap(),
);
tokio::time::sleep(Duration::from_millis(1000)).await;
scheduler.run_task_gen();
assert_eq!(scheduler.handles.len(), 1);
}
#[test]
fn test_task_builder_invalid_schedule() {
let result = TaskBuilder::new(chrono::Utc).every("* * * * * * *").build();
assert!(result.is_ok());
let result = TaskBuilder::new(chrono::Utc).every("invalid cron").build();
assert!(result.is_err());
match result {
Err(TaskError::InvalidCronExpression(_)) => {} _ => panic!("Expected InvalidCronExpression error"),
}
}
#[tokio::test]
async fn test_scheduler_graceful_shutdown() {
let mut scheduler = TaskScheduler::new(50, Local);
scheduler
.add_task(Task::new("* * * * * * *", None, None, Local))
.unwrap();
let handle = scheduler.handle();
let stopper = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(150)).await;
handle.shutdown();
});
tokio::time::timeout(Duration::from_secs(5), scheduler.run())
.await
.expect("scheduler did not stop after shutdown was requested");
stopper.await.unwrap();
assert_eq!(scheduler.handles.len(), 0);
}
#[tokio::test]
async fn test_scheduler_shutdown_requested_before_run() {
let mut scheduler = TaskScheduler::new(50, Local);
let handle = scheduler.handle();
handle.shutdown();
tokio::time::timeout(Duration::from_secs(5), scheduler.run())
.await
.expect("scheduler ignored a shutdown requested before run()");
}
#[tokio::test]
async fn test_scheduler_run_until() {
let mut scheduler = TaskScheduler::new(50, Local);
scheduler
.add_task(Task::new("* * * * * * *", None, None, Local))
.unwrap();
tokio::time::timeout(
Duration::from_secs(5),
scheduler.run_until(tokio::time::sleep(Duration::from_millis(120))),
)
.await
.expect("run_until did not stop when its shutdown future resolved");
assert_eq!(scheduler.handles.len(), 0);
}
}