Skip to main content

flare_core_runtime/state/
tracker.rs

1//! 状态追踪器实现
2
3use super::event::StateEvent;
4use crate::task::TaskState;
5use std::collections::HashMap;
6use std::sync::Arc;
7use std::time::{Duration, Instant};
8use tokio::sync::{RwLock, broadcast};
9
10/// 任务状态信息
11#[derive(Debug, Clone)]
12pub struct TaskStateInfo {
13    /// 任务名称
14    pub name: String,
15    /// 当前状态
16    pub state: TaskState,
17    /// 状态变更时间
18    pub state_changed_at: Instant,
19    /// 启动时间(如果已启动)
20    pub started_at: Option<Instant>,
21    /// 运行时长(如果正在运行)
22    pub running_duration: Option<Duration>,
23}
24
25impl TaskStateInfo {
26    /// 创建新的任务状态信息
27    pub fn new(name: impl Into<String>, initial_state: TaskState) -> Self {
28        Self {
29            name: name.into(),
30            state: initial_state,
31            state_changed_at: Instant::now(),
32            started_at: None,
33            running_duration: None,
34        }
35    }
36
37    /// 更新状态
38    pub fn update_state(&mut self, new_state: TaskState) {
39        let now = Instant::now();
40
41        // 记录启动时间
42        if new_state == TaskState::Running && self.started_at.is_none() {
43            self.started_at = Some(now);
44        }
45
46        // 计算运行时长
47        if new_state == TaskState::Running {
48            if let Some(started_at) = self.started_at {
49                self.running_duration = Some(now.duration_since(started_at));
50            }
51        } else {
52            self.running_duration = None;
53        }
54
55        self.state = new_state;
56        self.state_changed_at = now;
57    }
58}
59
60/// 状态追踪器
61///
62/// 实时追踪所有任务的状态,支持事件订阅
63///
64/// # 特性
65///
66/// - 线程安全(使用 `Arc<RwLock>`)
67/// - 支持事件订阅(使用 `broadcast` 通道)
68/// - 记录状态变更时间
69/// - 计算运行时长
70///
71/// # 示例
72///
73/// ```rust
74/// use flare_core_runtime::state::StateTracker;
75/// use flare_core_runtime::task::TaskState;
76///
77/// #[tokio::main]
78/// async fn main() {
79/// let tracker = StateTracker::new();
80///
81/// // 注册任务
82/// tracker.register_task("task-1", TaskState::Pending).await;
83///
84/// // 更新状态
85/// tracker.update_state("task-1", TaskState::Starting).await;
86/// tracker.update_state("task-1", TaskState::Running).await;
87///
88/// // 获取状态
89/// let info = tracker.get_state("task-1").await.unwrap();
90/// assert_eq!(info.state, TaskState::Running);
91/// }
92/// ```
93pub struct StateTracker {
94    /// 任务状态映射
95    tasks: Arc<RwLock<HashMap<String, TaskStateInfo>>>,
96    /// 事件发送器
97    event_tx: broadcast::Sender<StateEvent>,
98}
99
100impl StateTracker {
101    /// 创建新的状态追踪器
102    pub fn new() -> Self {
103        let (event_tx, _) = broadcast::channel(100);
104        Self {
105            tasks: Arc::new(RwLock::new(HashMap::new())),
106            event_tx,
107        }
108    }
109
110    /// 创建新的状态追踪器(指定事件缓冲区大小)
111    pub fn with_capacity(capacity: usize) -> Self {
112        let (event_tx, _) = broadcast::channel(capacity);
113        Self {
114            tasks: Arc::new(RwLock::new(HashMap::new())),
115            event_tx,
116        }
117    }
118
119    /// 注册任务
120    pub async fn register_task(&self, name: impl Into<String>, initial_state: TaskState) {
121        let name = name.into();
122        let info = TaskStateInfo::new(&name, initial_state);
123
124        let mut tasks = self.tasks.write().await;
125        tasks.insert(name, info);
126    }
127
128    /// 更新任务状态
129    ///
130    /// # 返回
131    ///
132    /// 返回状态事件(如果状态变更成功)
133    pub async fn update_state(&self, name: &str, new_state: TaskState) -> Option<StateEvent> {
134        let mut tasks = self.tasks.write().await;
135
136        if let Some(info) = tasks.get_mut(name) {
137            let old_state = info.state;
138
139            // 检查状态转换是否有效
140            if !old_state.can_transition_to(new_state) {
141                return None;
142            }
143
144            // 更新状态
145            info.update_state(new_state);
146
147            // 创建事件
148            let event = StateEvent::new(name, old_state, new_state);
149
150            // 发送事件(忽略错误,如果没有订阅者)
151            let _ = self.event_tx.send(event.clone());
152
153            return Some(event);
154        }
155
156        None
157    }
158
159    /// 更新任务状态(带错误信息)
160    pub async fn update_state_with_error(
161        &self,
162        name: &str,
163        new_state: TaskState,
164        error: impl Into<String>,
165    ) -> Option<StateEvent> {
166        let mut tasks = self.tasks.write().await;
167
168        if let Some(info) = tasks.get_mut(name) {
169            let old_state = info.state;
170
171            // 检查状态转换是否有效
172            if !old_state.can_transition_to(new_state) {
173                return None;
174            }
175
176            // 更新状态
177            info.update_state(new_state);
178
179            // 创建事件
180            let event = StateEvent::new(name, old_state, new_state).with_error(error);
181
182            // 发送事件
183            let _ = self.event_tx.send(event.clone());
184
185            return Some(event);
186        }
187
188        None
189    }
190
191    /// 获取任务状态
192    pub async fn get_state(&self, name: &str) -> Option<TaskStateInfo> {
193        let tasks = self.tasks.read().await;
194        tasks.get(name).cloned()
195    }
196
197    /// 获取所有任务状态
198    pub async fn get_all_states(&self) -> Vec<TaskStateInfo> {
199        let tasks = self.tasks.read().await;
200        tasks.values().cloned().collect()
201    }
202
203    /// 订阅状态事件
204    pub fn subscribe(&self) -> broadcast::Receiver<StateEvent> {
205        self.event_tx.subscribe()
206    }
207
208    /// 检查所有任务是否都已就绪
209    pub async fn all_ready(&self) -> bool {
210        let tasks = self.tasks.read().await;
211        tasks.values().all(|info| info.state == TaskState::Running)
212    }
213
214    /// 检查是否有任务失败
215    pub async fn has_failures(&self) -> bool {
216        let tasks = self.tasks.read().await;
217        tasks.values().any(|info| info.state == TaskState::Failed)
218    }
219
220    /// 获取失败的任务
221    pub async fn get_failed_tasks(&self) -> Vec<String> {
222        let tasks = self.tasks.read().await;
223        tasks
224            .values()
225            .filter(|info| info.state == TaskState::Failed)
226            .map(|info| info.name.clone())
227            .collect()
228    }
229
230    /// 获取运行中的任务数量
231    pub async fn running_count(&self) -> usize {
232        let tasks = self.tasks.read().await;
233        tasks
234            .values()
235            .filter(|info| info.state == TaskState::Running)
236            .count()
237    }
238}
239
240impl Default for StateTracker {
241    fn default() -> Self {
242        Self::new()
243    }
244}
245
246impl Clone for StateTracker {
247    fn clone(&self) -> Self {
248        Self {
249            tasks: Arc::clone(&self.tasks),
250            event_tx: self.event_tx.clone(),
251        }
252    }
253}
254
255#[cfg(test)]
256mod tests {
257    use super::*;
258
259    #[tokio::test]
260    async fn test_state_tracker_register() {
261        let tracker = StateTracker::new();
262        tracker.register_task("task-1", TaskState::Pending).await;
263
264        let info = tracker.get_state("task-1").await.unwrap();
265        assert_eq!(info.state, TaskState::Pending);
266    }
267
268    #[tokio::test]
269    async fn test_state_tracker_update() {
270        let tracker = StateTracker::new();
271        tracker.register_task("task-1", TaskState::Pending).await;
272
273        let event = tracker.update_state("task-1", TaskState::Starting).await;
274        assert!(event.is_some());
275
276        let info = tracker.get_state("task-1").await.unwrap();
277        assert_eq!(info.state, TaskState::Starting);
278    }
279
280    #[tokio::test]
281    async fn test_state_tracker_invalid_transition() {
282        let tracker = StateTracker::new();
283        tracker.register_task("task-1", TaskState::Pending).await;
284
285        // 无效转换:Pending -> Running(应该先 Starting)
286        let event = tracker.update_state("task-1", TaskState::Running).await;
287        assert!(event.is_none());
288    }
289
290    #[tokio::test]
291    async fn test_state_tracker_subscribe() {
292        let tracker = StateTracker::new();
293        let mut rx = tracker.subscribe();
294
295        tracker.register_task("task-1", TaskState::Pending).await;
296
297        // 在另一个任务中更新状态
298        let tracker_clone = tracker.clone();
299        tokio::spawn(async move {
300            tracker_clone
301                .update_state("task-1", TaskState::Starting)
302                .await;
303        });
304
305        // 接收事件
306        let event = rx.recv().await.unwrap();
307        assert_eq!(event.task_name, "task-1");
308        assert_eq!(event.old_state, TaskState::Pending);
309        assert_eq!(event.new_state, TaskState::Starting);
310    }
311
312    #[tokio::test]
313    async fn test_state_tracker_all_ready() {
314        let tracker = StateTracker::new();
315
316        tracker.register_task("task-1", TaskState::Pending).await;
317        tracker.register_task("task-2", TaskState::Pending).await;
318
319        assert!(!tracker.all_ready().await);
320
321        tracker.update_state("task-1", TaskState::Starting).await;
322        tracker.update_state("task-2", TaskState::Starting).await;
323        tracker.update_state("task-1", TaskState::Running).await;
324        tracker.update_state("task-2", TaskState::Running).await;
325
326        assert!(tracker.all_ready().await);
327    }
328}