flare_core_runtime/task/
manager.rs1use super::{Task, TaskResult, TaskState};
6use crate::config::RuntimeConfig;
7use crate::error::RuntimeError;
8use crate::state::StateTracker;
9use crate::utils::topological_sort;
10use std::collections::HashMap;
11use std::sync::Arc;
12use std::time::Duration;
13use tokio::sync::oneshot;
14use tokio::task::JoinSet;
15use tracing::{debug, error, info, warn};
16
17pub struct TaskManager {
49 tasks: Vec<Box<dyn Task>>,
51 state_tracker: Arc<StateTracker>,
53 config: RuntimeConfig,
55}
56
57impl TaskManager {
58 pub fn new() -> Self {
60 Self {
61 tasks: Vec::new(),
62 state_tracker: Arc::new(StateTracker::new()),
63 config: RuntimeConfig::default(),
64 }
65 }
66
67 pub fn with_config(config: RuntimeConfig) -> Self {
69 Self {
70 tasks: Vec::new(),
71 state_tracker: Arc::new(StateTracker::new()),
72 config,
73 }
74 }
75
76 pub(crate) fn set_config(&mut self, config: RuntimeConfig) {
77 self.config = config;
78 }
79
80 pub(crate) fn shutdown_timeout(&self) -> Duration {
81 self.config.shutdown_timeout
82 }
83
84 pub fn add_task(&mut self, task: Box<dyn Task>) {
86 debug!(task_name = %task.name(), "Adding task to manager");
87 self.tasks.push(task);
88 }
89
90 pub fn task_count(&self) -> usize {
92 self.tasks.len()
93 }
94
95 pub fn state_tracker(&self) -> Arc<StateTracker> {
97 Arc::clone(&self.state_tracker)
98 }
99
100 pub async fn start_all(
112 &mut self,
113 ) -> Result<(JoinSet<TaskResult>, Vec<oneshot::Sender<()>>), RuntimeError> {
114 info!(task_count = self.tasks.len(), "Starting all tasks");
115
116 let sorted_tasks = self.sort_tasks()?;
118
119 for task in &sorted_tasks {
121 self.state_tracker
122 .register_task(task.name(), TaskState::Pending)
123 .await;
124 }
125
126 let mut join_set = JoinSet::new();
128 let mut shutdown_txs = Vec::new();
129
130 for task in sorted_tasks {
131 let task_name = task.name().to_string();
132 let (shutdown_tx, shutdown_rx) = oneshot::channel();
133 shutdown_txs.push(shutdown_tx);
134
135 self.state_tracker
137 .update_state(&task_name, TaskState::Starting)
138 .await;
139
140 let state_tracker = Arc::clone(&self.state_tracker);
142 let task_future = task.run(shutdown_rx);
143
144 join_set.spawn(async move {
145 state_tracker
147 .update_state(&task_name, TaskState::Running)
148 .await;
149
150 let result = task_future.await;
152
153 match &result {
155 Ok(_) => {
156 debug!(task_name = %task_name, "Task completed");
157 state_tracker
158 .update_state(&task_name, TaskState::Stopped)
159 .await;
160 }
161 Err(e) => {
162 error!(task_name = %task_name, error = %e, "❌ Task failed");
163 state_tracker
164 .update_state_with_error(&task_name, TaskState::Failed, e.to_string())
165 .await;
166 }
167 }
168
169 result
170 });
171 }
172
173 info!("All tasks started");
174 Ok((join_set, shutdown_txs))
175 }
176
177 pub async fn stop_all(
184 &self,
185 mut join_set: JoinSet<TaskResult>,
186 shutdown_txs: Vec<oneshot::Sender<()>>,
187 ) {
188 info!("Stopping all tasks");
189
190 for tx in shutdown_txs {
192 let _ = tx.send(());
193 }
194
195 match tokio::time::timeout(self.config.shutdown_timeout, async {
197 while let Some(result) = join_set.join_next().await {
198 match result {
199 Ok(Ok(_)) => {
200 debug!("Task completed gracefully");
201 }
202 Ok(Err(e)) => {
203 warn!("Task completed with error: {}", e);
204 }
205 Err(e) => {
206 warn!("Task join error: {}", e);
207 }
208 }
209 }
210 })
211 .await
212 {
213 Ok(_) => {
214 info!("All tasks completed");
215 }
216 Err(_) => {
217 warn!("Tasks shutdown timeout, forcing exit");
218 join_set.abort_all();
219 }
220 }
221 }
222
223 pub async fn wait_for_ready(&self) -> Result<(), RuntimeError> {
225 if !self.config.task_startup.enable_ready_check {
226 debug!("Task ready check is disabled, skipping");
227 return Ok(());
228 }
229
230 info!("Waiting for all tasks to be ready");
231
232 let timeout = self.config.task_startup.ready_check_timeout;
234 let start = std::time::Instant::now();
235
236 loop {
237 if self.state_tracker.all_ready().await {
238 info!("✅ All tasks are ready");
239 return Ok(());
240 }
241
242 if start.elapsed() > timeout {
243 return Err(RuntimeError::StartupTimeout {
244 name: "all tasks".to_string(),
245 timeout,
246 });
247 }
248
249 tokio::time::sleep(Duration::from_millis(100)).await;
250 }
251 }
252
253 fn sort_tasks(&mut self) -> Result<Vec<Box<dyn Task>>, RuntimeError> {
255 let items: Vec<(String, Vec<String>)> = self
257 .tasks
258 .iter()
259 .map(|task| (task.name().to_string(), task.dependencies()))
260 .collect();
261
262 let sorted_names = topological_sort(items)
264 .map_err(|cycle| RuntimeError::CircularDependency { tasks: cycle })?;
265
266 let mut by_name: HashMap<String, Box<dyn Task>> = HashMap::new();
268 for task in self.tasks.drain(..) {
269 by_name.insert(task.name().to_string(), task);
270 }
271
272 let mut sorted_tasks = Vec::with_capacity(sorted_names.len());
273 for name in sorted_names {
274 if let Some(task) = by_name.remove(&name) {
275 sorted_tasks.push(task);
276 }
277 }
278 for (_, task) in by_name {
279 warn!(
280 task_name = %task.name(),
281 "Task missing from topological order output, appending at end"
282 );
283 sorted_tasks.push(task);
284 }
285
286 Ok(sorted_tasks)
287 }
288}
289
290impl Default for TaskManager {
291 fn default() -> Self {
292 Self::new()
293 }
294}
295
296#[cfg(test)]
297mod tests {
298 use super::*;
299 use crate::task::SpawnTask;
300
301 #[tokio::test]
302 async fn test_task_manager_new() {
303 let manager = TaskManager::new();
304 assert_eq!(manager.task_count(), 0);
305 }
306
307 #[tokio::test]
308 async fn test_task_manager_add_task() {
309 let mut manager = TaskManager::new();
310 manager.add_task(Box::new(SpawnTask::new("task-1", async { Ok(()) })));
311 assert_eq!(manager.task_count(), 1);
312 }
313
314 #[tokio::test]
315 async fn test_task_manager_state_tracker() {
316 let manager = TaskManager::new();
317 let _tracker = manager.state_tracker();
318 }
319
320 #[tokio::test]
322 async fn test_task_manager_start_all_three_independent_no_panic() {
323 let mut manager = TaskManager::new();
324 manager.add_task(Box::new(SpawnTask::new("conversation-grpc", async {
325 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
326 Ok(())
327 })));
328 manager.add_task(Box::new(SpawnTask::new("read-receipt-consumer", async {
329 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
330 Ok(())
331 })));
332 manager.add_task(Box::new(SpawnTask::new(
333 "conversation-ensure-consumer",
334 async {
335 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
336 Ok(())
337 },
338 )));
339
340 let (mut join_set, shutdown_txs) = manager.start_all().await.expect("start_all");
341 for tx in shutdown_txs {
342 let _ = tx.send(());
343 }
344 while join_set.join_next().await.is_some() {}
345 }
346}