mod config;
mod dependency_resolver;
mod running_tasks;
mod task_queue;
pub use config::SchedulerConfig;
pub use dependency_resolver::DependencyResolver;
pub use running_tasks::RunningTasks;
pub use task_queue::TaskQueue;
use crate::error::Result;
use crate::storage::TaskStorage;
use crate::task::{Task, TaskDefinition, TaskStatus};
use chrono::Utc;
use std::sync::Arc;
use tokio::time::{Duration, interval};
use tracing::{error, info, warn};
pub struct Scheduler {
storage: Arc<dyn TaskStorage>,
config: SchedulerConfig,
task_queue: TaskQueue,
running_tasks: RunningTasks,
dependency_resolver: DependencyResolver,
}
impl Scheduler {
pub fn new(storage: Arc<dyn TaskStorage>, config: SchedulerConfig) -> Self {
let storage_clone = Arc::clone(&storage);
Self {
storage: Arc::clone(&storage),
config,
task_queue: TaskQueue::new(),
running_tasks: RunningTasks::new(),
dependency_resolver: DependencyResolver::new(storage_clone),
}
}
pub async fn submit_task(&self, definition: TaskDefinition) -> Result<String> {
let task = Task::new(definition);
let task_id = task.definition.id.clone();
if self.config.enable_dependency_resolution {
self.dependency_resolver
.validate_dependencies(&task)
.await?;
}
self.storage.save_task(&task).await?;
info!("Task submitted: {} ({})", task.definition.name, task_id);
Ok(task_id)
}
pub async fn start(&self) -> Result<()> {
info!(
"Starting scheduler with poll interval: {}s",
self.config.poll_interval_seconds
);
let mut interval = interval(Duration::from_secs(self.config.poll_interval_seconds));
loop {
interval.tick().await;
if let Err(e) = self.schedule_cycle().await {
error!("Scheduler cycle failed: {}", e);
}
if let Err(e) = self.cleanup_old_tasks().await {
warn!("Task cleanup failed: {}", e);
}
}
}
async fn schedule_cycle(&self) -> Result<()> {
let running_count = self.running_tasks.len().await;
if running_count >= self.config.max_concurrent_tasks {
return Ok(());
}
let available_slots = self.config.max_concurrent_tasks - running_count;
let pending_tasks = self.storage.get_pending_tasks(available_slots).await?;
for task in pending_tasks {
if self.can_execute_task(&task).await? {
self.queue_task_for_execution(task).await?;
}
}
Ok(())
}
async fn can_execute_task(&self, task: &Task) -> Result<bool> {
if !task.is_ready_to_execute() {
return Ok(false);
}
if let Some(scheduled_at) = task.definition.scheduled_at {
if Utc::now() < scheduled_at {
return Ok(false);
}
}
if self.config.enable_dependency_resolution {
return self
.dependency_resolver
.are_dependencies_satisfied(task)
.await;
}
Ok(true)
}
async fn queue_task_for_execution(&self, mut task: Task) -> Result<()> {
let task_id = task.definition.id.clone();
task.status = TaskStatus::Running;
task.started_at = Some(Utc::now());
self.storage.save_task(&task).await?;
self.running_tasks.add(task_id.clone()).await;
self.task_queue.push(task).await;
info!("Task queued for execution: {}", task_id);
Ok(())
}
pub async fn get_next_task(&self) -> Option<Task> {
self.task_queue.pop().await
}
pub async fn complete_task(
&self,
task_id: &str,
success: bool,
output: Option<String>,
error: Option<String>,
) -> Result<()> {
if let Some(mut task) = self.storage.get_task(task_id).await? {
let execution_time = task
.started_at
.map(|start| (Utc::now() - start).num_milliseconds() as u64)
.unwrap_or(0);
let result = crate::task::TaskResult {
success,
output,
error,
execution_time_ms: execution_time,
metadata: std::collections::HashMap::new(),
};
task.complete_execution(result);
self.storage.save_task(&task).await?;
self.running_tasks.remove(task_id).await;
if !success && task.can_retry() {
self.schedule_retry(&task).await?;
}
info!("Task completed: {} (success: {})", task_id, success);
}
Ok(())
}
async fn schedule_retry(&self, task: &Task) -> Result<()> {
let mut retry_task = task.clone();
retry_task.retry();
let delay_seconds = self.calculate_retry_delay(retry_task.retry_count);
let scheduled_at = Utc::now() + chrono::Duration::seconds(delay_seconds as i64);
retry_task.definition.scheduled_at = Some(scheduled_at);
self.storage.save_task(&retry_task).await?;
info!(
"Task scheduled for retry: {} (attempt: {})",
retry_task.definition.id, retry_task.retry_count
);
Ok(())
}
fn calculate_retry_delay(&self, retry_count: u32) -> u64 {
std::cmp::min(2_u64.pow(retry_count) * 30, 300)
}
async fn cleanup_old_tasks(&self) -> Result<()> {
let cutoff_time = Utc::now()
- chrono::Duration::hours(self.config.cleanup_completed_tasks_after_hours as i64);
let completed_tasks = self
.storage
.list_tasks_by_status(TaskStatus::Completed)
.await?;
let failed_tasks = self
.storage
.list_tasks_by_status(TaskStatus::Failed)
.await?;
let mut cleanup_count = 0;
for task in completed_tasks.into_iter().chain(failed_tasks.into_iter()) {
if let Some(completed_at) = task.completed_at {
if completed_at < cutoff_time {
self.storage.delete_task(&task.definition.id).await?;
cleanup_count += 1;
}
}
}
if cleanup_count > 0 {
info!("Cleaned up {} old tasks", cleanup_count);
}
Ok(())
}
pub async fn cancel_task(&self, task_id: &str) -> Result<()> {
if let Some(mut task) = self.storage.get_task(task_id).await? {
task.cancel();
self.storage.save_task(&task).await?;
self.running_tasks.remove(task_id).await;
info!("Task cancelled: {}", task_id);
}
Ok(())
}
pub async fn get_task_status(&self, task_id: &str) -> Result<Option<TaskStatus>> {
if let Some(task) = self.storage.get_task(task_id).await? {
Ok(Some(task.status))
} else {
Ok(None)
}
}
pub async fn list_tasks(&self, status: Option<TaskStatus>) -> Result<Vec<Task>> {
match status {
Some(s) => self.storage.list_tasks_by_status(s).await,
None => {
let mut all_tasks = Vec::new();
for status in [
TaskStatus::Pending,
TaskStatus::Running,
TaskStatus::Completed,
TaskStatus::Failed,
TaskStatus::Cancelled,
] {
let mut tasks = self.storage.list_tasks_by_status(status).await?;
all_tasks.append(&mut tasks);
}
Ok(all_tasks)
}
}
}
}