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