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    #[cfg(test)]
81    pub(crate) fn shutdown_timeout(&self) -> Duration {
82        self.config.shutdown_timeout
83    }
84
85    /// 添加任务
86    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    /// 获取任务数量
92    pub fn task_count(&self) -> usize {
93        self.tasks.len()
94    }
95
96    /// 获取状态追踪器
97    pub fn state_tracker(&self) -> Arc<StateTracker> {
98        Arc::clone(&self.state_tracker)
99    }
100
101    /// 启动所有任务
102    ///
103    /// # 返回
104    ///
105    /// - `join_set` - 任务 JoinSet,用于等待任务完成
106    /// - `shutdown_txs` - shutdown 信号发送器列表
107    ///
108    /// # 错误
109    ///
110    /// - 循环依赖
111    /// - 缺失依赖
112    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        // 1. 拓扑排序,确定启动顺序
118        let sorted_tasks = self.sort_tasks()?;
119
120        // 2. 注册所有任务到状态追踪器
121        for task in &sorted_tasks {
122            self.state_tracker
123                .register_task(task.name(), TaskState::Pending)
124                .await;
125        }
126
127        // 3. 启动任务
128        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            // 更新状态为 Starting
137            self.state_tracker
138                .update_state(&task_name, TaskState::Starting)
139                .await;
140
141            // 启动任务
142            let state_tracker = Arc::clone(&self.state_tracker);
143            let task_future = task.run(shutdown_rx);
144
145            join_set.spawn(async move {
146                // 更新状态为 Running
147                state_tracker
148                    .update_state(&task_name, TaskState::Running)
149                    .await;
150
151                // 执行任务
152                let result = task_future.await;
153
154                // 更新状态
155                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    /// 停止所有任务
179    ///
180    /// # 参数
181    ///
182    /// * `join_set` - 任务 JoinSet
183    /// * `shutdown_txs` - shutdown 信号发送器列表
184    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        // 1. 发送 shutdown 信号
192        for tx in shutdown_txs {
193            let _ = tx.send(());
194        }
195
196        // 2. 等待所有任务关闭
197        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    /// 等待所有任务就绪
225    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        // 等待所有任务状态变为 Running
234        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    /// 拓扑排序任务
255    fn sort_tasks(&mut self) -> Result<Vec<Box<dyn Task>>, RuntimeError> {
256        // 构建任务依赖列表
257        let items: Vec<(String, Vec<String>)> = self
258            .tasks
259            .iter()
260            .map(|task| (task.name().to_string(), task.dependencies()))
261            .collect();
262
263        // 拓扑排序
264        let sorted_names = topological_sort(items)
265            .map_err(|cycle| RuntimeError::CircularDependency { tasks: cycle })?;
266
267        // 按拓扑名顺序取出任务:禁止用「原始下标 + swap_remove」——每次 swap_remove 都会打乱下标,导致越界 panic。
268        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    /// 回归:多任务且无依赖时,拓扑序与插入序可能不一致,旧实现用 swap_remove(原下标) 会 panic。
322    #[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}