use super::event::StateEvent;
use crate::task::TaskState;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{RwLock, broadcast};
#[derive(Debug, Clone)]
pub struct TaskStateInfo {
pub name: String,
pub state: TaskState,
pub state_changed_at: Instant,
pub started_at: Option<Instant>,
pub running_duration: Option<Duration>,
}
impl TaskStateInfo {
pub fn new(name: impl Into<String>, initial_state: TaskState) -> Self {
Self {
name: name.into(),
state: initial_state,
state_changed_at: Instant::now(),
started_at: None,
running_duration: None,
}
}
pub fn update_state(&mut self, new_state: TaskState) {
let now = Instant::now();
if new_state == TaskState::Running && self.started_at.is_none() {
self.started_at = Some(now);
}
if new_state == TaskState::Running {
if let Some(started_at) = self.started_at {
self.running_duration = Some(now.duration_since(started_at));
}
} else {
self.running_duration = None;
}
self.state = new_state;
self.state_changed_at = now;
}
}
pub struct StateTracker {
tasks: Arc<RwLock<HashMap<String, TaskStateInfo>>>,
event_tx: broadcast::Sender<StateEvent>,
}
impl StateTracker {
pub fn new() -> Self {
let (event_tx, _) = broadcast::channel(100);
Self {
tasks: Arc::new(RwLock::new(HashMap::new())),
event_tx,
}
}
pub fn with_capacity(capacity: usize) -> Self {
let (event_tx, _) = broadcast::channel(capacity);
Self {
tasks: Arc::new(RwLock::new(HashMap::new())),
event_tx,
}
}
pub async fn register_task(&self, name: impl Into<String>, initial_state: TaskState) {
let name = name.into();
let info = TaskStateInfo::new(&name, initial_state);
let mut tasks = self.tasks.write().await;
tasks.insert(name, info);
}
pub async fn update_state(&self, name: &str, new_state: TaskState) -> Option<StateEvent> {
let mut tasks = self.tasks.write().await;
if let Some(info) = tasks.get_mut(name) {
let old_state = info.state;
if !old_state.can_transition_to(new_state) {
return None;
}
info.update_state(new_state);
let event = StateEvent::new(name, old_state, new_state);
let _ = self.event_tx.send(event.clone());
return Some(event);
}
None
}
pub async fn update_state_with_error(
&self,
name: &str,
new_state: TaskState,
error: impl Into<String>,
) -> Option<StateEvent> {
let mut tasks = self.tasks.write().await;
if let Some(info) = tasks.get_mut(name) {
let old_state = info.state;
if !old_state.can_transition_to(new_state) {
return None;
}
info.update_state(new_state);
let event = StateEvent::new(name, old_state, new_state).with_error(error);
let _ = self.event_tx.send(event.clone());
return Some(event);
}
None
}
pub async fn get_state(&self, name: &str) -> Option<TaskStateInfo> {
let tasks = self.tasks.read().await;
tasks.get(name).cloned()
}
pub async fn get_all_states(&self) -> Vec<TaskStateInfo> {
let tasks = self.tasks.read().await;
tasks.values().cloned().collect()
}
pub fn subscribe(&self) -> broadcast::Receiver<StateEvent> {
self.event_tx.subscribe()
}
pub async fn all_ready(&self) -> bool {
let tasks = self.tasks.read().await;
tasks.values().all(|info| info.state == TaskState::Running)
}
pub async fn has_failures(&self) -> bool {
let tasks = self.tasks.read().await;
tasks.values().any(|info| info.state == TaskState::Failed)
}
pub async fn get_failed_tasks(&self) -> Vec<String> {
let tasks = self.tasks.read().await;
tasks
.values()
.filter(|info| info.state == TaskState::Failed)
.map(|info| info.name.clone())
.collect()
}
pub async fn running_count(&self) -> usize {
let tasks = self.tasks.read().await;
tasks
.values()
.filter(|info| info.state == TaskState::Running)
.count()
}
}
impl Default for StateTracker {
fn default() -> Self {
Self::new()
}
}
impl Clone for StateTracker {
fn clone(&self) -> Self {
Self {
tasks: Arc::clone(&self.tasks),
event_tx: self.event_tx.clone(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_state_tracker_register() {
let tracker = StateTracker::new();
tracker.register_task("task-1", TaskState::Pending).await;
let info = tracker.get_state("task-1").await.unwrap();
assert_eq!(info.state, TaskState::Pending);
}
#[tokio::test]
async fn test_state_tracker_update() {
let tracker = StateTracker::new();
tracker.register_task("task-1", TaskState::Pending).await;
let event = tracker.update_state("task-1", TaskState::Starting).await;
assert!(event.is_some());
let info = tracker.get_state("task-1").await.unwrap();
assert_eq!(info.state, TaskState::Starting);
}
#[tokio::test]
async fn test_state_tracker_invalid_transition() {
let tracker = StateTracker::new();
tracker.register_task("task-1", TaskState::Pending).await;
let event = tracker.update_state("task-1", TaskState::Running).await;
assert!(event.is_none());
}
#[tokio::test]
async fn test_state_tracker_subscribe() {
let tracker = StateTracker::new();
let mut rx = tracker.subscribe();
tracker.register_task("task-1", TaskState::Pending).await;
let tracker_clone = tracker.clone();
tokio::spawn(async move {
tracker_clone
.update_state("task-1", TaskState::Starting)
.await;
});
let event = rx.recv().await.unwrap();
assert_eq!(event.task_name, "task-1");
assert_eq!(event.old_state, TaskState::Pending);
assert_eq!(event.new_state, TaskState::Starting);
}
#[tokio::test]
async fn test_state_tracker_all_ready() {
let tracker = StateTracker::new();
tracker.register_task("task-1", TaskState::Pending).await;
tracker.register_task("task-2", TaskState::Pending).await;
assert!(!tracker.all_ready().await);
tracker.update_state("task-1", TaskState::Starting).await;
tracker.update_state("task-2", TaskState::Starting).await;
tracker.update_state("task-1", TaskState::Running).await;
tracker.update_state("task-2", TaskState::Running).await;
assert!(tracker.all_ready().await);
}
}