Skip to main content

taskflow_rs/scheduler/
mod.rs

1mod config;
2mod dependency_resolver;
3mod running_tasks;
4mod task_queue;
5
6pub use config::SchedulerConfig;
7pub use dependency_resolver::DependencyResolver;
8pub use running_tasks::RunningTasks;
9pub use task_queue::TaskQueue;
10
11use crate::error::Result;
12use crate::storage::TaskStorage;
13use crate::task::{Task, TaskDefinition, TaskStatus};
14use chrono::Utc;
15use std::sync::Arc;
16use tokio::time::{Duration, interval};
17use tracing::{error, info, warn};
18
19pub struct Scheduler {
20    storage: Arc<dyn TaskStorage>,
21    config: SchedulerConfig,
22    task_queue: TaskQueue,
23    running_tasks: RunningTasks,
24    dependency_resolver: DependencyResolver,
25}
26
27impl Scheduler {
28    pub fn new(storage: Arc<dyn TaskStorage>, config: SchedulerConfig) -> Self {
29        let storage_clone = Arc::clone(&storage);
30        Self {
31            storage: Arc::clone(&storage),
32            config,
33            task_queue: TaskQueue::new(),
34            running_tasks: RunningTasks::new(),
35            dependency_resolver: DependencyResolver::new(storage_clone),
36        }
37    }
38
39    pub async fn submit_task(&self, definition: TaskDefinition) -> Result<String> {
40        let task = Task::new(definition);
41        let task_id = task.definition.id.clone();
42
43        if self.config.enable_dependency_resolution {
44            self.dependency_resolver
45                .validate_dependencies(&task)
46                .await?;
47        }
48
49        self.storage.save_task(&task).await?;
50
51        info!("Task submitted: {} ({})", task.definition.name, task_id);
52        Ok(task_id)
53    }
54
55    pub async fn start(&self) -> Result<()> {
56        info!(
57            "Starting scheduler with poll interval: {}s",
58            self.config.poll_interval_seconds
59        );
60
61        let mut interval = interval(Duration::from_secs(self.config.poll_interval_seconds));
62
63        loop {
64            interval.tick().await;
65
66            if let Err(e) = self.schedule_cycle().await {
67                error!("Scheduler cycle failed: {}", e);
68            }
69
70            if let Err(e) = self.cleanup_old_tasks().await {
71                warn!("Task cleanup failed: {}", e);
72            }
73        }
74    }
75
76    async fn schedule_cycle(&self) -> Result<()> {
77        let running_count = self.running_tasks.len().await;
78        if running_count >= self.config.max_concurrent_tasks {
79            return Ok(());
80        }
81
82        let available_slots = self.config.max_concurrent_tasks - running_count;
83        let pending_tasks = self.storage.get_pending_tasks(available_slots).await?;
84
85        for task in pending_tasks {
86            if self.can_execute_task(&task).await? {
87                self.queue_task_for_execution(task).await?;
88            }
89        }
90
91        Ok(())
92    }
93
94    async fn can_execute_task(&self, task: &Task) -> Result<bool> {
95        if !task.is_ready_to_execute() {
96            return Ok(false);
97        }
98
99        if let Some(scheduled_at) = task.definition.scheduled_at {
100            if Utc::now() < scheduled_at {
101                return Ok(false);
102            }
103        }
104
105        if self.config.enable_dependency_resolution {
106            return self
107                .dependency_resolver
108                .are_dependencies_satisfied(task)
109                .await;
110        }
111
112        Ok(true)
113    }
114
115    async fn queue_task_for_execution(&self, mut task: Task) -> Result<()> {
116        let task_id = task.definition.id.clone();
117
118        task.status = TaskStatus::Running;
119        task.started_at = Some(Utc::now());
120        self.storage.save_task(&task).await?;
121
122        self.running_tasks.add(task_id.clone()).await;
123        self.task_queue.push(task).await;
124
125        info!("Task queued for execution: {}", task_id);
126        Ok(())
127    }
128
129    pub async fn get_next_task(&self) -> Option<Task> {
130        self.task_queue.pop().await
131    }
132
133    pub async fn complete_task(
134        &self,
135        task_id: &str,
136        success: bool,
137        output: Option<String>,
138        error: Option<String>,
139    ) -> Result<()> {
140        if let Some(mut task) = self.storage.get_task(task_id).await? {
141            let execution_time = task
142                .started_at
143                .map(|start| (Utc::now() - start).num_milliseconds() as u64)
144                .unwrap_or(0);
145
146            let result = crate::task::TaskResult {
147                success,
148                output,
149                error,
150                execution_time_ms: execution_time,
151                metadata: std::collections::HashMap::new(),
152            };
153
154            task.complete_execution(result);
155            self.storage.save_task(&task).await?;
156
157            self.running_tasks.remove(task_id).await;
158
159            if !success && task.can_retry() {
160                self.schedule_retry(&task).await?;
161            }
162
163            info!("Task completed: {} (success: {})", task_id, success);
164        }
165
166        Ok(())
167    }
168
169    async fn schedule_retry(&self, task: &Task) -> Result<()> {
170        let mut retry_task = task.clone();
171        retry_task.retry();
172
173        let delay_seconds = self.calculate_retry_delay(retry_task.retry_count);
174        let scheduled_at = Utc::now() + chrono::Duration::seconds(delay_seconds as i64);
175        retry_task.definition.scheduled_at = Some(scheduled_at);
176
177        self.storage.save_task(&retry_task).await?;
178
179        info!(
180            "Task scheduled for retry: {} (attempt: {})",
181            retry_task.definition.id, retry_task.retry_count
182        );
183
184        Ok(())
185    }
186
187    fn calculate_retry_delay(&self, retry_count: u32) -> u64 {
188        std::cmp::min(2_u64.pow(retry_count) * 30, 300)
189    }
190
191    async fn cleanup_old_tasks(&self) -> Result<()> {
192        let cutoff_time = Utc::now()
193            - chrono::Duration::hours(self.config.cleanup_completed_tasks_after_hours as i64);
194
195        let completed_tasks = self
196            .storage
197            .list_tasks_by_status(TaskStatus::Completed)
198            .await?;
199        let failed_tasks = self
200            .storage
201            .list_tasks_by_status(TaskStatus::Failed)
202            .await?;
203
204        let mut cleanup_count = 0;
205
206        for task in completed_tasks.into_iter().chain(failed_tasks.into_iter()) {
207            if let Some(completed_at) = task.completed_at {
208                if completed_at < cutoff_time {
209                    self.storage.delete_task(&task.definition.id).await?;
210                    cleanup_count += 1;
211                }
212            }
213        }
214
215        if cleanup_count > 0 {
216            info!("Cleaned up {} old tasks", cleanup_count);
217        }
218
219        Ok(())
220    }
221
222    pub async fn cancel_task(&self, task_id: &str) -> Result<()> {
223        if let Some(mut task) = self.storage.get_task(task_id).await? {
224            task.cancel();
225            self.storage.save_task(&task).await?;
226
227            self.running_tasks.remove(task_id).await;
228
229            info!("Task cancelled: {}", task_id);
230        }
231        Ok(())
232    }
233
234    pub async fn get_task_status(&self, task_id: &str) -> Result<Option<TaskStatus>> {
235        if let Some(task) = self.storage.get_task(task_id).await? {
236            Ok(Some(task.status))
237        } else {
238            Ok(None)
239        }
240    }
241
242    pub async fn list_tasks(&self, status: Option<TaskStatus>) -> Result<Vec<Task>> {
243        match status {
244            Some(s) => self.storage.list_tasks_by_status(s).await,
245            None => {
246                let mut all_tasks = Vec::new();
247                for status in [
248                    TaskStatus::Pending,
249                    TaskStatus::Running,
250                    TaskStatus::Completed,
251                    TaskStatus::Failed,
252                    TaskStatus::Cancelled,
253                ] {
254                    let mut tasks = self.storage.list_tasks_by_status(status).await?;
255                    all_tasks.append(&mut tasks);
256                }
257                Ok(all_tasks)
258            }
259        }
260    }
261}