Skip to main content

flare_core_runtime/task/
manager.rs

1//! 任务管理器实现
2//!
3//! 负责管理所有任务的生命周期,包括启动、停止、依赖管理等
4
5use 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
17/// 任务管理器
18///
19/// 负责管理所有任务的生命周期
20///
21/// # 功能
22///
23/// - 任务注册和管理
24/// - 依赖排序和启动
25/// - 优雅停机
26/// - 状态追踪
27///
28/// # 示例
29///
30/// ```rust,ignore
31/// use flare_core_runtime::task::TaskManager;
32/// use flare_core_runtime::task::SpawnTask;
33///
34/// let manager = TaskManager::new();
35///
36/// // 添加任务
37/// manager.add_task(Box::new(SpawnTask::new("task-1", async { Ok(()) })));
38///
39/// // 启动所有任务
40/// let (join_set, shutdown_txs) = manager.start_all().await?;
41///
42/// // 等待停机信号
43/// // ...
44///
45/// // 停止所有任务
46/// manager.stop_all(join_set, shutdown_txs).await;
47/// ```
48pub struct TaskManager {
49    /// 任务列表
50    tasks: Vec<Box<dyn Task>>,
51    /// 状态追踪器
52    state_tracker: Arc<StateTracker>,
53    /// 配置
54    config: RuntimeConfig,
55}
56
57impl TaskManager {
58    /// 创建新的任务管理器
59    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    /// 创建新的任务管理器(带配置)
68    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    /// 添加任务
85    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    /// 获取任务数量
91    pub fn task_count(&self) -> usize {
92        self.tasks.len()
93    }
94
95    /// 获取状态追踪器
96    pub fn state_tracker(&self) -> Arc<StateTracker> {
97        Arc::clone(&self.state_tracker)
98    }
99
100    /// 启动所有任务
101    ///
102    /// # 返回
103    ///
104    /// - `join_set` - 任务 JoinSet,用于等待任务完成
105    /// - `shutdown_txs` - shutdown 信号发送器列表
106    ///
107    /// # 错误
108    ///
109    /// - 循环依赖
110    /// - 缺失依赖
111    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        // 1. 拓扑排序,确定启动顺序
117        let sorted_tasks = self.sort_tasks()?;
118
119        // 2. 注册所有任务到状态追踪器
120        for task in &sorted_tasks {
121            self.state_tracker
122                .register_task(task.name(), TaskState::Pending)
123                .await;
124        }
125
126        // 3. 启动任务
127        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            // 更新状态为 Starting
136            self.state_tracker
137                .update_state(&task_name, TaskState::Starting)
138                .await;
139
140            // 启动任务
141            let state_tracker = Arc::clone(&self.state_tracker);
142            let task_future = task.run(shutdown_rx);
143
144            join_set.spawn(async move {
145                // 更新状态为 Running
146                state_tracker
147                    .update_state(&task_name, TaskState::Running)
148                    .await;
149
150                // 执行任务
151                let result = task_future.await;
152
153                // 更新状态
154                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    /// 停止所有任务
178    ///
179    /// # 参数
180    ///
181    /// * `join_set` - 任务 JoinSet
182    /// * `shutdown_txs` - shutdown 信号发送器列表
183    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        // 1. 发送 shutdown 信号
191        for tx in shutdown_txs {
192            let _ = tx.send(());
193        }
194
195        // 2. 等待所有任务关闭
196        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    /// 等待所有任务就绪
224    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        // 等待所有任务状态变为 Running
233        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    /// 拓扑排序任务
254    fn sort_tasks(&mut self) -> Result<Vec<Box<dyn Task>>, RuntimeError> {
255        // 构建任务依赖列表
256        let items: Vec<(String, Vec<String>)> = self
257            .tasks
258            .iter()
259            .map(|task| (task.name().to_string(), task.dependencies()))
260            .collect();
261
262        // 拓扑排序
263        let sorted_names = topological_sort(items)
264            .map_err(|cycle| RuntimeError::CircularDependency { tasks: cycle })?;
265
266        // 按拓扑名顺序取出任务:禁止用「原始下标 + swap_remove」——每次 swap_remove 都会打乱下标,导致越界 panic。
267        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    /// 回归:多任务且无依赖时,拓扑序与插入序可能不一致,旧实现用 swap_remove(原下标) 会 panic。
321    #[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}