use std::sync::Mutex;
use nautilus_core::MUTEX_POISONED;
use super::dst::task::JoinHandle;
#[derive(Debug, Default)]
pub struct TaskHandles {
handles: Mutex<Vec<JoinHandle<()>>>,
}
impl TaskHandles {
pub fn push(&self, handle: JoinHandle<()>) {
let mut handles = self.handles.lock().expect(MUTEX_POISONED);
handles.retain(|handle| !handle.is_finished());
handles.push(handle);
}
pub fn abort_all(&self) {
for handle in self.take_all() {
handle.abort();
}
}
#[must_use]
pub fn take_all(&self) -> Vec<JoinHandle<()>> {
let mut handles = self.handles.lock().expect(MUTEX_POISONED);
std::mem::take(&mut *handles)
}
#[must_use]
pub fn all_finished(&self) -> bool {
self.handles
.lock()
.expect(MUTEX_POISONED)
.iter()
.all(JoinHandle::is_finished)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.handles.lock().expect(MUTEX_POISONED).is_empty()
}
#[must_use]
pub fn len(&self) -> usize {
self.handles.lock().expect(MUTEX_POISONED).len()
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use rstest::rstest;
use super::*;
use crate::live::dst::{task, time};
#[rstest]
#[cfg_attr(not(all(feature = "simulation", madsim)), tokio::test)]
#[cfg_attr(all(feature = "simulation", madsim), madsim::test)]
async fn test_push_prunes_finished_handles() {
let tasks = TaskHandles::default();
let finished = task::spawn(async {});
time::timeout(Duration::from_secs(1), async {
while !finished.is_finished() {
task::yield_now().await;
}
})
.await
.expect("task should finish");
tasks.push(finished);
tasks.push(task::spawn(std::future::pending()));
assert_eq!(tasks.len(), 1);
tasks.abort_all();
}
#[rstest]
#[cfg_attr(not(all(feature = "simulation", madsim)), tokio::test)]
#[cfg_attr(all(feature = "simulation", madsim), madsim::test)]
async fn test_abort_all_drains_before_aborting() {
let tasks = TaskHandles::default();
let mut drop_receivers = Vec::new();
for _ in 0..2 {
let (drop_tx, drop_rx) = tokio::sync::oneshot::channel();
let signal = DropSignal { tx: Some(drop_tx) };
tasks.push(task::spawn(async move {
let _signal = signal;
std::future::pending::<()>().await;
}));
drop_receivers.push(drop_rx);
}
tasks.abort_all();
assert!(tasks.is_empty());
for drop_rx in drop_receivers {
time::timeout(Duration::from_secs(1), drop_rx)
.await
.expect("aborted task should drop its future")
.expect("drop signal should be sent");
}
}
#[rstest]
#[cfg_attr(not(all(feature = "simulation", madsim)), tokio::test)]
#[cfg_attr(all(feature = "simulation", madsim), madsim::test)]
async fn test_take_all_extracts_handles() {
let tasks = TaskHandles::default();
tasks.push(task::spawn(std::future::pending()));
tasks.push(task::spawn(std::future::pending()));
let handles = tasks.take_all();
assert!(tasks.is_empty());
assert_eq!(handles.len(), 2);
for handle in handles {
handle.abort();
}
}
#[rstest]
#[cfg_attr(not(all(feature = "simulation", madsim)), tokio::test)]
#[cfg_attr(all(feature = "simulation", madsim), madsim::test)]
async fn test_all_finished_preserves_handles() {
let tasks = TaskHandles::default();
assert!(tasks.all_finished());
let (release_tx, release_rx) = tokio::sync::oneshot::channel();
tasks.push(task::spawn(async move {
let _ = release_rx.await;
}));
assert!(!tasks.all_finished());
assert_eq!(tasks.len(), 1);
release_tx.send(()).expect("task should still be waiting");
time::timeout(Duration::from_secs(1), async {
while !tasks.all_finished() {
task::yield_now().await;
}
})
.await
.expect("task should finish");
assert!(tasks.all_finished());
assert_eq!(tasks.len(), 1);
}
#[rstest]
#[cfg_attr(not(all(feature = "simulation", madsim)), tokio::test)]
#[cfg_attr(all(feature = "simulation", madsim), madsim::test)]
async fn test_drop_detaches_tasks() {
let tasks = TaskHandles::default();
let (release_tx, release_rx) = tokio::sync::oneshot::channel();
let (done_tx, done_rx) = tokio::sync::oneshot::channel();
tasks.push(task::spawn(async move {
let _ = release_rx.await;
let _ = done_tx.send(());
}));
drop(tasks);
let _ = release_tx.send(());
time::timeout(Duration::from_secs(1), done_rx)
.await
.expect("detached task should complete")
.expect("completion signal should be sent");
}
struct DropSignal {
tx: Option<tokio::sync::oneshot::Sender<()>>,
}
impl Drop for DropSignal {
fn drop(&mut self) {
if let Some(tx) = self.tx.take() {
let _ = tx.send(());
}
}
}
}