flare_core_runtime/state/
tracker.rs1use 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#[derive(Debug, Clone)]
12pub struct TaskStateInfo {
13 pub name: String,
15 pub state: TaskState,
17 pub state_changed_at: Instant,
19 pub started_at: Option<Instant>,
21 pub running_duration: Option<Duration>,
23}
24
25impl TaskStateInfo {
26 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 pub fn update_state(&mut self, new_state: TaskState) {
39 let now = Instant::now();
40
41 if new_state == TaskState::Running && self.started_at.is_none() {
43 self.started_at = Some(now);
44 }
45
46 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
60pub struct StateTracker {
94 tasks: Arc<RwLock<HashMap<String, TaskStateInfo>>>,
96 event_tx: broadcast::Sender<StateEvent>,
98}
99
100impl StateTracker {
101 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 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 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 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 if !old_state.can_transition_to(new_state) {
141 return None;
142 }
143
144 info.update_state(new_state);
146
147 let event = StateEvent::new(name, old_state, new_state);
149
150 let _ = self.event_tx.send(event.clone());
152
153 return Some(event);
154 }
155
156 None
157 }
158
159 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 if !old_state.can_transition_to(new_state) {
173 return None;
174 }
175
176 info.update_state(new_state);
178
179 let event = StateEvent::new(name, old_state, new_state).with_error(error);
181
182 let _ = self.event_tx.send(event.clone());
184
185 return Some(event);
186 }
187
188 None
189 }
190
191 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 pub async fn get_all_states(&self) -> Vec<TaskStateInfo> {
199 let tasks = self.tasks.read().await;
200 tasks.values().cloned().collect()
201 }
202
203 pub fn subscribe(&self) -> broadcast::Receiver<StateEvent> {
205 self.event_tx.subscribe()
206 }
207
208 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 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 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 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 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 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 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}