You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在递归任务中使用线程池?(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();
}

关键逻辑说明

  1. 原子计数器:Arc<AtomicUsize>用于跨线程共享并安全修改任务计数,每提交一个任务计数器加1,任务完成(包括递归子任务)时减1。
  2. 主线程等待:主线程循环检查计数器,直到归0,此时所有任务(包括递归生成的子任务)都已执行完毕。
  3. 线程池关闭:计数器归0后,主线程调用shutdown()关闭任务发送端,线程池中的线程会在接收不到任务时自动退出,主线程等待所有线程结束。

注意事项

  • 使用crossbeam-channel的无界通道,比标准库的std::sync::mpsc更适合高并发场景(也可替换为标准库通道)。
  • 任务中的计数器操作必须保证原子性,使用SeqCst内存序确保计数的准确性。
  • 主线程的等待逻辑用yield_now()减少CPU占用,也可使用条件变量优化(比如Arc<Condvar>配合Mutex),避免忙等。

内容的提问来源于Stack Exchange,提问作者Saurabh Goyal

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.22 15:54:55