use std::collections::HashMap;
use std::future::Future;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Instant;
use tokio::sync::Mutex;
use tokio::task::{AbortHandle, JoinHandle};
pub type TaskId = u64;
#[derive(Debug, Clone)]
#[allow(dead_code)] pub struct TaskInfo {
pub id: TaskId,
pub name: String,
pub created_at: Instant,
}
#[allow(dead_code)] struct TaskEntry {
info: TaskInfo,
abort_handle: AbortHandle,
}
pub struct TaskRegistry {
next_id: AtomicU64,
tasks: Arc<Mutex<HashMap<TaskId, TaskEntry>>>,
}
impl TaskRegistry {
pub fn new() -> Self {
TaskRegistry {
next_id: AtomicU64::new(1),
tasks: Arc::new(Mutex::new(HashMap::new())),
}
}
fn generate_id(&self) -> TaskId {
self.next_id.fetch_add(1, Ordering::SeqCst)
}
pub async fn spawn_tracked<F>(&self, name: impl Into<String>, future: F) -> TaskId
where
F: Future<Output = ()> + Send + 'static,
{
let id = self.generate_id();
let name = name.into();
let tasks = self.tasks.clone();
let tasks_cleanup = self.tasks.clone();
let join_handle: JoinHandle<()> = tokio::spawn(async move {
future.await;
tasks_cleanup.lock().await.remove(&id);
});
let abort_handle = join_handle.abort_handle();
let entry = TaskEntry {
info: TaskInfo {
id,
name,
created_at: Instant::now(),
},
abort_handle,
};
tasks.lock().await.insert(id, entry);
id
}
#[allow(dead_code)] pub async fn cancel(&self, task_id: TaskId) -> bool {
let mut tasks = self.tasks.lock().await;
if let Some(entry) = tasks.remove(&task_id) {
entry.abort_handle.abort();
true
} else {
false
}
}
#[allow(dead_code)] pub async fn get_active_tasks(&self) -> Vec<TaskInfo> {
let tasks = self.tasks.lock().await;
tasks.values().map(|e| e.info.clone()).collect()
}
#[allow(dead_code)] pub async fn active_count(&self) -> usize {
self.tasks.lock().await.len()
}
#[allow(dead_code)] pub async fn is_active(&self, task_id: TaskId) -> bool {
self.tasks.lock().await.contains_key(&task_id)
}
#[allow(dead_code)] pub async fn cancel_all(&self) {
let mut tasks = self.tasks.lock().await;
for entry in tasks.values() {
entry.abort_handle.abort();
}
tasks.clear();
}
#[allow(dead_code)] pub async fn cleanup_finished(&self) {
let mut tasks = self.tasks.lock().await;
tasks.retain(|_, entry| !entry.abort_handle.is_finished());
}
}
impl Default for TaskRegistry {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicBool;
use tokio::time::{sleep, Duration};
#[tokio::test]
async fn test_spawn_tracked_task() {
let registry = TaskRegistry::new();
let task_id = registry
.spawn_tracked("test-task", async {
sleep(Duration::from_millis(50)).await;
})
.await;
assert!(task_id > 0);
assert!(registry.is_active(task_id).await);
sleep(Duration::from_millis(100)).await;
registry.cleanup_finished().await;
assert!(!registry.is_active(task_id).await);
}
#[tokio::test]
async fn test_task_cancellation() {
let registry = TaskRegistry::new();
let was_cancelled = Arc::new(AtomicBool::new(false));
let was_cancelled_clone = was_cancelled.clone();
let task_id = registry
.spawn_tracked("long-task", async move {
sleep(Duration::from_secs(10)).await;
was_cancelled_clone.store(true, Ordering::SeqCst);
})
.await;
assert!(registry.is_active(task_id).await);
let cancelled = registry.cancel(task_id).await;
assert!(cancelled);
sleep(Duration::from_millis(10)).await;
assert!(!registry.is_active(task_id).await);
assert!(!was_cancelled.load(Ordering::SeqCst));
}
#[tokio::test]
async fn test_get_active_tasks() {
let registry = TaskRegistry::new();
let _id1 = registry
.spawn_tracked("task-1", async {
sleep(Duration::from_secs(10)).await;
})
.await;
let _id2 = registry
.spawn_tracked("task-2", async {
sleep(Duration::from_secs(10)).await;
})
.await;
let active = registry.get_active_tasks().await;
assert_eq!(active.len(), 2);
let names: Vec<_> = active.iter().map(|t| t.name.as_str()).collect();
assert!(names.contains(&"task-1"));
assert!(names.contains(&"task-2"));
registry.cancel_all().await;
}
#[tokio::test]
async fn test_cancel_all() {
let registry = TaskRegistry::new();
for i in 0..5 {
registry
.spawn_tracked(format!("task-{}", i), async {
sleep(Duration::from_secs(10)).await;
})
.await;
}
assert_eq!(registry.active_count().await, 5);
registry.cancel_all().await;
assert_eq!(registry.active_count().await, 0);
}
#[tokio::test]
async fn test_task_registry_cleanup() {
let registry = TaskRegistry::new();
let task_id = registry
.spawn_tracked("quick-task", async {
sleep(Duration::from_millis(10)).await;
})
.await;
assert!(registry.is_active(task_id).await);
sleep(Duration::from_millis(50)).await;
registry.cleanup_finished().await;
assert!(!registry.is_active(task_id).await);
assert_eq!(registry.active_count().await, 0);
}
#[tokio::test]
async fn test_task_id_generation() {
let registry = TaskRegistry::new();
let id1 = registry
.spawn_tracked("task-1", async {})
.await;
let id2 = registry
.spawn_tracked("task-2", async {})
.await;
let id3 = registry
.spawn_tracked("task-3", async {})
.await;
assert!(id1 < id2);
assert!(id2 < id3);
}
}