Rayon线程池退出回调未执行及任务同步问题排查
Rayon线程池任务执行与线程退出问题排查
问题场景
我正在学习Rayon线程池,希望实现一个程序:向线程池添加一批任务,当该批次最后一个任务执行时添加更多任务。编写的Rust代码如下:
use rayon; use std::sync::atomic::AtomicU32; use std::sync::{Arc, Mutex}; fn main() { let thread_num = std::thread::available_parallelism().unwrap().into(); let pool = rayon::ThreadPoolBuilder::new() .num_threads(thread_num) .exit_handler(|i| { println!("Thread {} exited", i); }) .build() .unwrap(); let pool = Arc::new(pool); let counter = Arc::new(AtomicU32::new(thread_num as u32)); let test = Arc::new(Mutex::new(String::default())); for i in 0..thread_num { let _pool = pool.clone(); let _counter = counter.clone(); let _test = test.clone(); pool.spawn(move || { let val = _counter.fetch_sub(1, std::sync::atomic::Ordering::SeqCst); _test .lock() .unwrap() .push_str(format!("b {}\n", val).as_str()); if val == 1 { _counter.store(thread_num as u32, std::sync::atomic::Ordering::SeqCst); for i in 0..thread_num { let _test = _test.clone(); _pool.spawn(move || { _test.lock().unwrap().push_str("aaaa\n"); }); } } }); } println!("{}", test.lock().unwrap().as_str()); }
运行后出现两个问题:
- 输出结果不稳定,有时包含"aaaa"有时没有,说明打印时部分任务尚未执行完成;
- 线程退出回调函数未被调用。
错误原因分析
1. 输出不稳定:主线程未等待任务完成
主线程提交第一批任务后直接执行println!,但Rayon线程池的任务是后台异步执行的,此时新增的"aaaa"任务可能还未完成,导致打印结果不确定。Rayon不会自动阻塞主线程等待所有任务结束,必须手动同步。
2. 线程退出回调未触发
Rayon线程池的线程会保持存活直到线程池实例被彻底销毁(所有Arc引用被drop)。主线程结束时,虽然Arc<ThreadPool>会被drop,但主线程退出速度太快,线程池来不及优雅关闭线程,因此退出回调无法被调用。
修正方案与代码
核心修正点
- 添加原子计数器跟踪所有待完成任务,让主线程等待任务全部执行完毕;
- 手动销毁线程池引用并短暂等待,确保线程池优雅关闭,触发退出回调。
修正后的代码:
use rayon; use std::sync::atomic::{AtomicU32, Ordering}; use std::sync::{Arc, Mutex}; use std::thread; fn main() { let thread_num = thread::available_parallelism().unwrap().into(); let pool = rayon::ThreadPoolBuilder::new() .num_threads(thread_num) .exit_handler(|i| { println!("Thread {} exited", i); }) .build() .unwrap(); let pool = Arc::new(pool); let counter = Arc::new(AtomicU32::new(thread_num as u32)); let test = Arc::new(Mutex::new(String::default())); // 跟踪所有待完成的任务总数 let total_tasks = Arc::new(AtomicU32::new(thread_num as u32)); for i in 0..thread_num { let _pool = pool.clone(); let _counter = counter.clone(); let _test = test.clone(); let _total_tasks = total_tasks.clone(); pool.spawn(move || { let val = _counter.fetch_sub(1, Ordering::SeqCst); _test.lock().unwrap().push_str(format!("b {}\n", val).as_str()); if val == 1 { _counter.store(thread_num as u32, Ordering::SeqCst); // 添加新任务时更新总任务数 _total_tasks.fetch_add(thread_num as u32, Ordering::SeqCst); for i in 0..thread_num { let _test = _test.clone(); let _total_tasks_inner = _total_tasks.clone(); _pool.spawn(move || { _test.lock().unwrap().push_str("aaaa\n"); // 任务完成后递减计数器 _total_tasks_inner.fetch_sub(1, Ordering::SeqCst); }); } } // 当前任务完成后递减计数器 _total_tasks.fetch_sub(1, Ordering::SeqCst); }); } // 主线程等待所有任务执行完毕 while total_tasks.load(Ordering::SeqCst) > 0 { thread::yield_now(); } // 打印最终结果 println!("{}", test.lock().unwrap().as_str()); // 销毁线程池的Arc引用,触发线程池关闭 drop(pool); // 短暂等待线程池优雅退出,确保回调被调用 thread::sleep(std::time::Duration::from_millis(100)); }
内容的提问来源于stack exchange,提问作者Makogan
相关产品推荐
相关产品推荐

