//! Lexically scoped execution for finite blocking work. use tokio::runtime::{Handle, RuntimeFlavor}; /// Runs finite blocking work without detaching it from the calling task. /// /// On a multi-thread Tokio runtime, `block_in_place` lets the executor hand the /// caller's other asynchronous work to another worker while this closure runs. /// A current-thread runtime cannot make that handoff, so it executes the closure /// directly, as does code running outside a Tokio runtime. /// /// This function deliberately has no cancellation point: when it returns, the /// closure has completed and all values owned by the closure have been dropped. pub fn scoped_blocking(work: F) -> R where F: FnOnce() -> R, { if matches!( Handle::try_current().map(|handle| handle.runtime_flavor()), Ok(RuntimeFlavor::MultiThread) ) { tokio::task::block_in_place(work) } else { work() } } #[cfg(test)] mod tests { use std::{ sync::{Arc, Condvar, Mutex, mpsc}, thread, time::Duration, }; use tokio::runtime::{Handle, RuntimeFlavor}; use super::scoped_blocking; struct DropSignal(mpsc::Sender<()>); impl Drop for DropSignal { fn drop(&mut self) { let _ = self.0.send(()); } } #[test] fn runs_directly_without_a_runtime() { let calling_thread = thread::current().id(); let execution_thread = scoped_blocking(|| thread::current().id()); assert_eq!(execution_thread, calling_thread); } #[tokio::test] async fn runs_directly_on_a_current_thread_runtime() { assert_eq!( Handle::current().runtime_flavor(), RuntimeFlavor::CurrentThread ); let calling_thread = thread::current().id(); let execution_thread = scoped_blocking(|| thread::current().id()); assert_eq!(execution_thread, calling_thread); } #[tokio::test(flavor = "multi_thread", worker_threads = 1)] async fn abort_waits_for_in_place_work_to_finish() { assert_eq!( Handle::current().runtime_flavor(), RuntimeFlavor::MultiThread ); let gate = Arc::new((Mutex::new(false), Condvar::new())); let gate_for_task = Arc::clone(&gate); let gate_for_release_thread = Arc::clone(&gate); let (entered_tx, entered_rx) = tokio::sync::oneshot::channel(); let (dropped_tx, dropped_rx) = mpsc::channel(); let (request_release, release_requested) = mpsc::channel(); // The timeout also makes this test fail safe: if an assertion panics, // dropping `request_release` wakes this thread so the runtime can shut down. let release_thread = thread::spawn(move || { let _ = release_requested.recv_timeout(Duration::from_secs(2)); let (gate_open, wake) = &*gate_for_release_thread; let mut gate_open = gate_open.lock().expect("release gate must not be poisoned"); *gate_open = true; wake.notify_one(); }); let mut task = tokio::spawn(async move { scoped_blocking(move || { let _drop_signal = DropSignal(dropped_tx); entered_tx .send(()) .expect("test task must report entering blocking work"); let (gate_open, wake) = &*gate_for_task; let gate_open = gate_open.lock().expect("task gate must not be poisoned"); let _gate_open = wake .wait_while(gate_open, |gate_open| !*gate_open) .expect("task gate must not be poisoned"); }); }); tokio::time::timeout(Duration::from_secs(2), entered_rx) .await .expect("blocking work must start") .expect("blocking task must retain the entry sender"); task.abort(); assert!( tokio::time::timeout(Duration::from_millis(50), &mut task) .await .is_err(), "aborting the task must not complete it while in-place work is blocked" ); assert_eq!(dropped_rx.try_recv(), Err(mpsc::TryRecvError::Empty)); request_release .send(()) .expect("release thread must remain available"); release_thread .join() .expect("release thread must not panic"); let completion = tokio::time::timeout(Duration::from_secs(2), &mut task) .await .expect("task must finish after its blocking work is released"); match completion { Ok(()) => {} Err(error) if error.is_cancelled() => {} Err(error) => panic!("blocking task failed unexpectedly: {error}"), } dropped_rx .recv_timeout(Duration::from_secs(1)) .expect("blocking closure values must be dropped before task completion"); } }