如何在递归任务中使用线程池?(Rust文件遍历场景)
支持递归任务的线程池终止方案
问题背景
我需要对目录中的每个文件执行操作(比如打印文件名),单线程实现如下:
use std::{env, fs, path::Path, rc::Rc}; fn main() { let args = env::args().collect::<Vec<String>>(); let process: ProcessFunc = Rc::new(|p| println!("{p}")); process_file_path(String::from("/playground/target/release/deps"), Rc::clone(&process)); } type ProcessFunc = Rc<dyn Fn(String) + Send + 'static>; fn process_file_path(source: String, f: ProcessFunc) { let source_path = Path::new(&source) .canonicalize() .expect("invalid source path"); let source_path_str = source_path.as_path().to_str().unwrap().to_string(); f(source_path_str); if source_path.is_dir() { for child in fs::read_dir(source_path).unwrap() { let child = child.unwrap(); process_file_path( child.path().as_path().to_str().unwrap().to_string(), Rc::clone(&f), ) } } }
因为文件间无依赖,计划用线程池并发处理,但遇到线程终止难题:每个任务(路径)可能生成更多子任务(子路径),线程无法判断何时停止。之前尝试的两种方法都有问题:
- 给线程加超时:不可靠且增加等待时间
- 无任务时关闭线程:会导致处理完叶子节点后多数线程提前关闭,无法完成剩余任务
尝试过用两个channel让主线程判断任务状态,但仍无法确定所有任务完成的时机,线程池实现如下(未在main中使用):
struct ThreadPool<T> { job_sender: Option<Sender<Box<dyn FnOnce() -> T + Send + 'static>>>, handles: Vec<JoinHandle<()>>, } impl<T> ThreadPool<T> where T: Send + 'static, { fn new(cap: usize) -> (Self, Receiver<T>) { assert_ne!(cap, 0); let (job_sender, job_receiver) = channel::<Box<dyn FnOnce() -> T + Send + 'static>>(); let (result_sender, result_receiver) = channel::<T>(); let result_sender = Arc::new(Mutex::new(result_sender)); let job_receiver = Arc::new(Mutex::new(job_receiver)); let mut handles = vec![]; for _i in 0..cap { let job_receiver = Arc::clone(&job_receiver); let result_sender = Arc::clone(&result_sender); handles.push(thread::spawn(move || { while let Ok(job) = job_receiver.lock().unwrap().recv() { result_sender.lock().unwrap().send(job()).unwrap(); } })); } ( ThreadPool { handles, job_sender: Some(job_sender), }, result_receiver, ) } fn add(&self, job: Box<dyn FnOnce() -> T + Send + 'static>) { self.job_sender.as_ref().unwrap().send(job).unwrap(); } } impl<T> Drop for ThreadPool<T> { fn drop(&mut self) { self.job_sender = None; while let Some(handle) = self.handles.pop() { handle.join().unwrap(); } } }
解决方案:原子计数器跟踪任务生命周期
核心是通过原子计数器跟踪所有活跃任务(包括递归生成的子任务),确保主线程能准确判断所有任务完成的时机,再安全终止线程池。
完整实现代码:
use std::{ sync::{Arc, AtomicUsize, Mutex}, thread::{self, JoinHandle}, path::Path, fs, }; use crossbeam_channel::{unbounded, Sender, Receiver}; // 线程池结构体 struct ThreadPool { job_sender: Option<Sender<Box<dyn FnOnce() + Send + 'static>>>, handles: Vec<JoinHandle<()>>, } impl ThreadPool { fn new(cap: usize) -> Self { assert_ne!(cap, 0); let (job_sender, job_receiver) = unbounded::<Box<dyn FnOnce() + Send + 'static>>(); let job_receiver = Arc::new(Mutex::new(job_receiver)); let mut handles = vec![]; for _ in 0..cap { let job_receiver = Arc::clone(&job_receiver); handles.push(thread::spawn(move || { loop { // 尝试接收任务,若发送端已关闭则退出循环 let job = match job_receiver.lock().unwrap().recv() { Ok(job) => job, Err(_) => break, }; // 执行任务 job(); } })); } ThreadPool { job_sender: Some(job_sender), handles, } } // 提交任务,同时更新计数器 fn submit(&self, job: Box<dyn FnOnce() + Send + 'static>) { self.job_sender.as_ref().unwrap().send(job).unwrap(); } // 关闭任务发送端,等待所有线程退出 fn shutdown(&mut self) { self.job_sender.take(); for handle in self.handles.drain(..) { handle.join().unwrap(); } } } // 处理文件路径的核心逻辑,带任务计数 fn process_path( path_str: String, process_func: impl Fn(String) + Send + 'static, pool: &ThreadPool, task_counter: Arc<AtomicUsize>, ) { // 执行文件处理逻辑 let path = Path::new(&path_str).canonicalize().expect("invalid path"); let path_display = path.to_str().unwrap().to_string(); process_func(path_display.clone()); // 如果是目录,递归提交子路径任务 if path.is_dir() { for entry in fs::read_dir(path).unwrap() { let entry = entry.unwrap(); let child_path = entry.path().to_str().unwrap().to_string(); // 计数器加1(因为要提交新任务) task_counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst); // 克隆计数器和pool引用,提交任务 let counter_clone = Arc::clone(&task_counter); let pool_clone = pool; let func_clone = process_func.clone(); pool.submit(Box::new(move || { process_path(child_path, func_clone, pool_clone, counter_clone); // 任务完成,计数器减1 counter_clone.fetch_sub(1, std::sync::atomic::Ordering::SeqCst); })); } } // 当前任务完成,计数器减1 task_counter.fetch_sub(1, std::sync::atomic::Ordering::SeqCst); } fn main() { // 创建线程池(比如4个线程) let mut pool = ThreadPool::new(4); // 初始化任务计数器,初始值为1(第一个任务) let task_counter = Arc::new(AtomicUsize::new(1)); // 定义处理函数 let process_func = |path: String| println!("Processing: {}", path); // 提交根目录任务 let root_path = "/playground/target/release/deps".to_string(); let counter_clone = Arc::clone(&task_counter); pool.submit(Box::new(move || { process_path(root_path, process_func, &pool, counter_clone); })); // 主线程等待计数器归0,确保所有任务完成 while task_counter.load(std::sync::atomic::Ordering::SeqCst) != 0 { thread::yield_now(); // 让出CPU,避免忙等 } // 关闭线程池 pool.shutdown(); }
关键逻辑说明
- 原子计数器:
Arc<AtomicUsize>用于跨线程共享并安全修改任务计数,每提交一个任务计数器加1,任务完成(包括递归子任务)时减1。 - 主线程等待:主线程循环检查计数器,直到归0,此时所有任务(包括递归生成的子任务)都已执行完毕。
- 线程池关闭:计数器归0后,主线程调用
shutdown()关闭任务发送端,线程池中的线程会在接收不到任务时自动退出,主线程等待所有线程结束。
注意事项
- 使用
crossbeam-channel的无界通道,比标准库的std::sync::mpsc更适合高并发场景(也可替换为标准库通道)。 - 任务中的计数器操作必须保证原子性,使用
SeqCst内存序确保计数的准确性。 - 主线程的等待逻辑用
yield_now()减少CPU占用,也可使用条件变量优化(比如Arc<Condvar>配合Mutex),避免忙等。
内容的提问来源于Stack Exchange,提问作者Saurabh Goyal
相关产品推荐
相关产品推荐

