use std::future::Future; use futures::{StreamExt, stream::FuturesUnordered}; use tokio_util::sync::CancellationToken; /// Collects structured child futures, draining them before reporting cancellation. /// /// Children share `cancel_token` and are responsible for cooperatively settling /// their own resources. Keeping them unspawned also guarantees that dropping the /// collector synchronously drops every child instead of detaching background work. pub(super) async fn collect_or_drain_on_cancel( mut in_flight: FuturesUnordered, cancel_token: &CancellationToken, game_id: &str, ) -> eyre::Result> where F: Future, { let mut results = Vec::with_capacity(in_flight.len()); let mut cancelled = cancel_token.is_cancelled(); while !in_flight.is_empty() { if cancelled { let _ = in_flight .next() .await .expect("in-flight future should exist"); continue; } tokio::select! { biased; () = cancel_token.cancelled() => { cancelled = true; results.clear(); } result = in_flight.next() => { results.push(result.expect("in-flight future should exist")); } } } if cancelled { eyre::bail!("download cancelled for game {game_id}"); } Ok(results) } #[cfg(test)] mod tests { use std::{ future::Future, pin::Pin, sync::{ Arc, atomic::{AtomicUsize, Ordering}, }, time::Duration, }; use futures::stream::FuturesUnordered; use tokio::sync::{Notify, Semaphore, mpsc}; use tokio_util::sync::CancellationToken; use super::collect_or_drain_on_cancel; type Worker = Pin + Send>>; struct DropProbe(Arc); impl Drop for DropProbe { fn drop(&mut self) { self.0.fetch_add(1, Ordering::SeqCst); } } #[tokio::test] async fn cancellation_waits_for_every_child_to_quiesce() { let cancel_token = CancellationToken::new(); let started = Arc::new(AtomicUsize::new(0)); let completed = Arc::new(AtomicUsize::new(0)); let started_notify = Arc::new(Notify::new()); let release = Arc::new(Semaphore::new(0)); let workers = FuturesUnordered::::new(); for _ in 0..2 { let started = Arc::clone(&started); let completed = Arc::clone(&completed); let started_notify = Arc::clone(&started_notify); let release = Arc::clone(&release); workers.push(Box::pin(async move { started.fetch_add(1, Ordering::SeqCst); started_notify.notify_one(); release .acquire_owned() .await .expect("test semaphore should remain open") .forget(); completed.fetch_add(1, Ordering::SeqCst); })); } let collector_token = cancel_token.clone(); let mut collector = tokio::spawn(async move { collect_or_drain_on_cancel(workers, &collector_token, "game").await }); while started.load(Ordering::SeqCst) != 2 { started_notify.notified().await; } cancel_token.cancel(); assert!( tokio::time::timeout(Duration::from_millis(50), &mut collector) .await .is_err(), "cancellation must wait for child cleanup" ); release.add_permits(2); let error = collector .await .expect("collector task should join") .expect_err("cancelled collection should fail after draining"); assert!( error .to_string() .contains("download cancelled for game game") ); assert_eq!(completed.load(Ordering::SeqCst), 2); } #[tokio::test] async fn dropping_collector_drops_every_child() { let cancel_token = CancellationToken::new(); let dropped = Arc::new(AtomicUsize::new(0)); let (started_tx, mut started_rx) = mpsc::unbounded_channel(); let workers = FuturesUnordered::::new(); for worker_id in 0..2 { let probe = DropProbe(dropped.clone()); let worker_started = started_tx.clone(); workers.push(Box::pin(async move { let _probe = probe; worker_started .send(worker_id) .expect("collector should wait for workers"); std::future::pending::<()>().await; })); } drop(started_tx); let collector_token = cancel_token.clone(); let collector = tokio::spawn(async move { collect_or_drain_on_cancel(workers, &collector_token, "game").await }); started_rx.recv().await.expect("first worker should start"); started_rx.recv().await.expect("second worker should start"); collector.abort(); let join_error = collector .await .expect_err("collector should have been cancelled"); assert!(join_error.is_cancelled()); assert_eq!(dropped.load(Ordering::SeqCst), 2); } }