use std::sync::Arc;
use std::time::Duration;
use tokio_util::sync::CancellationToken;
use sz_orm_scheduler::{CronScheduler, JobHandler, ScheduledTask, Scheduler, SchedulerError};
#[derive(Debug, Clone)]
pub struct SchedulerRuntimeConfig {
pub tick_ms: u64,
}
impl Default for SchedulerRuntimeConfig {
fn default() -> Self {
Self { tick_ms: 1000 }
}
}
impl SchedulerRuntimeConfig {
pub fn new(tick_ms: u64) -> Self {
Self {
tick_ms: tick_ms.max(1),
}
}
}
pub struct SchedulerRuntime {
config: SchedulerRuntimeConfig,
scheduler: Arc<CronScheduler>,
}
impl SchedulerRuntime {
pub fn new(config: SchedulerRuntimeConfig) -> Self {
Self {
config,
scheduler: Arc::new(CronScheduler::new()),
}
}
pub fn start(&self, token: CancellationToken) -> tokio::task::JoinHandle<()> {
let scheduler = self.scheduler.clone();
let tick_interval = Duration::from_millis(self.config.tick_ms);
tokio::spawn(async move {
let mut ticker = tokio::time::interval(tick_interval);
loop {
tokio::select! {
_ = token.cancelled() => break,
_ = ticker.tick() => {
let now = chrono::Utc::now();
let fired = scheduler.try_fire_due(now);
if fired > 0 {
tracing::debug!("scheduler fired {} task(s) at {}", fired, now);
}
}
}
}
})
}
pub fn schedule(&self, task: ScheduledTask) -> Result<(), SchedulerError> {
self.scheduler.schedule(task)
}
pub fn cancel(&self, task_id: &str) -> Result<(), SchedulerError> {
self.scheduler.cancel(task_id)
}
pub fn pause(&self, task_id: &str) -> Result<(), SchedulerError> {
self.scheduler.pause(task_id)
}
pub fn resume(&self, task_id: &str) -> Result<(), SchedulerError> {
self.scheduler.resume(task_id)
}
pub fn list_tasks(&self) -> Vec<ScheduledTask> {
self.scheduler.list_tasks()
}
pub fn register_handler(&self, task_id: impl Into<String>, handler: Arc<dyn JobHandler>) {
self.scheduler.register_handler(task_id, handler);
}
pub fn try_fire_due(&self) -> usize {
let now = chrono::Utc::now();
self.scheduler.try_fire_due(now)
}
pub fn task_count(&self) -> usize {
self.scheduler.list_tasks().len()
}
pub fn config(&self) -> &SchedulerRuntimeConfig {
&self.config
}
pub fn scheduler(&self) -> &CronScheduler {
&self.scheduler
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
struct CounterHandler {
counter: Arc<AtomicUsize>,
}
impl CounterHandler {
fn new() -> (Self, Arc<AtomicUsize>) {
let counter = Arc::new(AtomicUsize::new(0));
let handler = Self {
counter: counter.clone(),
};
(handler, counter)
}
}
impl JobHandler for CounterHandler {
fn handle(&self, _task: &ScheduledTask) -> Result<(), String> {
self.counter.fetch_add(1, Ordering::SeqCst);
Ok(())
}
}
#[test]
fn test_scheduler_runtime_config_default() {
let config = SchedulerRuntimeConfig::default();
assert_eq!(config.tick_ms, 1000);
}
#[test]
fn test_scheduler_runtime_config_custom() {
let config = SchedulerRuntimeConfig::new(500);
assert_eq!(config.tick_ms, 500);
}
#[test]
fn test_scheduler_runtime_config_zero_clamped() {
let config = SchedulerRuntimeConfig::new(0);
assert_eq!(config.tick_ms, 1);
}
#[test]
fn test_schedule_task() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
let task = ScheduledTask::new("task-1", "测试任务", "0 * * * *");
runtime.schedule(task).unwrap();
assert_eq!(runtime.task_count(), 1);
}
#[test]
fn test_schedule_multiple_tasks() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
runtime
.schedule(ScheduledTask::new("t1", "任务1", "0 * * * *"))
.unwrap();
runtime
.schedule(ScheduledTask::new("t2", "任务2", "0 0 * * *"))
.unwrap();
runtime
.schedule(ScheduledTask::new("t3", "任务3", "0 0 0 * *"))
.unwrap();
assert_eq!(runtime.task_count(), 3);
}
#[test]
fn test_cancel_task() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
runtime
.schedule(ScheduledTask::new("task-1", "测试", "0 * * * *"))
.unwrap();
assert_eq!(runtime.task_count(), 1);
runtime.cancel("task-1").unwrap();
assert_eq!(runtime.task_count(), 0);
}
#[test]
fn test_cancel_nonexistent_task() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
let result = runtime.cancel("nonexistent");
assert!(result.is_err());
}
#[test]
fn test_pause_resume_task() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
runtime
.schedule(ScheduledTask::new("task-1", "测试", "0 * * * *"))
.unwrap();
runtime.pause("task-1").unwrap();
let tasks = runtime.list_tasks();
assert!(!tasks[0].enabled);
runtime.resume("task-1").unwrap();
let tasks = runtime.list_tasks();
assert!(tasks[0].enabled);
}
#[test]
fn test_list_tasks() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
runtime
.schedule(ScheduledTask::new("t1", "任务1", "0 * * * *"))
.unwrap();
runtime
.schedule(ScheduledTask::new("t2", "任务2", "0 0 * * *"))
.unwrap();
let tasks = runtime.list_tasks();
assert_eq!(tasks.len(), 2);
}
#[test]
fn test_register_handler() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
runtime
.schedule(ScheduledTask::new("task-1", "测试", "0 * * * *"))
.unwrap();
let (handler, _counter) = CounterHandler::new();
runtime.register_handler("task-1", Arc::new(handler));
}
#[tokio::test]
async fn test_start_and_cancel() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::new(10));
let token = CancellationToken::new();
let handle = runtime.start(token.clone());
tokio::time::sleep(Duration::from_millis(50)).await;
token.cancel();
let result = tokio::time::timeout(Duration::from_secs(2), handle).await;
assert!(result.is_ok(), "scheduler task should stop on cancel");
}
#[tokio::test]
async fn test_scheduler_fires_due_task() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::new(10));
runtime
.schedule(ScheduledTask::new("every-second", "每秒任务", "* * * * *"))
.unwrap();
let (handler, counter) = CounterHandler::new();
runtime.register_handler("every-second", Arc::new(handler));
let token = CancellationToken::new();
let handle = runtime.start(token.clone());
tokio::time::sleep(Duration::from_millis(100)).await;
token.cancel();
let _ = handle.await;
assert!(
counter.load(Ordering::SeqCst) >= 1,
"task should have fired at least once"
);
}
#[test]
fn test_try_fire_due_manual() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
runtime
.schedule(ScheduledTask::new("every-second", "每秒任务", "* * * * *"))
.unwrap();
let (handler, counter) = CounterHandler::new();
runtime.register_handler("every-second", Arc::new(handler));
let fired = runtime.try_fire_due();
assert!(fired >= 1);
assert!(counter.load(Ordering::SeqCst) >= 1);
}
#[test]
fn test_try_fire_due_no_tasks() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
let fired = runtime.try_fire_due();
assert_eq!(fired, 0);
}
#[test]
fn test_config_accessor() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::new(250));
assert_eq!(runtime.config().tick_ms, 250);
}
#[test]
fn test_scheduler_accessor() {
let runtime = SchedulerRuntime::new(SchedulerRuntimeConfig::default());
let _scheduler = runtime.scheduler();
}
}