Skip to main content

weavatrix_scan/parallel/
visit.rs

1use super::collect::{DirectoryTask, collect_shallow};
2use super::visit_worker::{visit_lane, visit_serial};
3use super::{ParallelWalker, parallel_worker_count};
4use crate::control::CancellationToken;
5use crate::pool::ThreadPool;
6use crate::walker::{ErrorPolicy, WalkEntry, WalkError};
7use std::collections::HashSet;
8use std::sync::{
9    Arc,
10    atomic::{AtomicBool, Ordering},
11    mpsc,
12};
13
14/// A streaming event emitted from a parallel traversal worker.
15#[derive(Debug)]
16pub enum WalkEvent<'a> {
17    Entry(&'a WalkEntry),
18    Error(&'a WalkError),
19}
20
21/// Controls traversal after a streaming visitor handles an event.
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum WalkControl {
24    Continue,
25    /// Prevents descent when the current event is a directory entry.
26    Skip,
27    /// Cooperatively stops every traversal worker.
28    Quit,
29}
30
31/// Summary of a streaming parallel traversal.
32#[derive(Debug)]
33pub struct ParallelVisitReport {
34    pub visited: u64,
35    pub errors: Vec<WalkError>,
36    pub quit: bool,
37    pub cancelled: bool,
38}
39
40impl ParallelWalker {
41    /// Visits entries directly on traversal workers without collecting paths.
42    ///
43    /// Visitor calls may run concurrently and their order is intentionally
44    /// unspecified. Use [`Self::walk`] when deterministic collected order is
45    /// required.
46    ///
47    /// # Errors
48    ///
49    /// Returns a root error or the first traversal error under
50    /// `ErrorPolicy::Abort`.
51    ///
52    /// # Panics
53    ///
54    /// Panics if the visitor or an internal traversal worker panics.
55    pub fn visit<F>(self, visitor: F) -> Result<ParallelVisitReport, WalkError>
56    where
57        F: for<'entry> Fn(WalkEvent<'entry>) -> WalkControl + Send + Sync + 'static,
58    {
59        self.visit_with_cancellation(&CancellationToken::new(), visitor)
60    }
61
62    /// Streaming parallel traversal with cooperative cancellation.
63    ///
64    /// # Errors
65    ///
66    /// Returns a root error or the first traversal error under
67    /// `ErrorPolicy::Abort`.
68    ///
69    /// # Panics
70    ///
71    /// Panics if the visitor or an internal traversal worker panics.
72    pub fn visit_with_cancellation<F>(
73        mut self,
74        cancellation: &CancellationToken,
75        visitor: F,
76    ) -> Result<ParallelVisitReport, WalkError>
77    where
78        F: for<'entry> Fn(WalkEvent<'entry>) -> WalkControl + Send + Sync + 'static,
79    {
80        self.options = self.options.normalized();
81        if self.options.follow_links {
82            return visit_serial(&self.root, self.options, cancellation, visitor);
83        }
84        let mut shallow = collect_shallow(&self.root, self.options)?;
85        let stop = Arc::new(AtomicBool::new(false));
86        let visitor = Arc::new(visitor);
87        let mut visited = 0_u64;
88        let mut errors = Vec::new();
89        let mut skipped = HashSet::new();
90        for entry in &shallow.entries {
91            if cancellation.is_cancelled() || stop.load(Ordering::Acquire) {
92                break;
93            }
94            visited = visited.saturating_add(1);
95            match visitor(WalkEvent::Entry(entry)) {
96                WalkControl::Skip if entry.depth() == 0 => {
97                    stop.store(true, Ordering::Release);
98                }
99                WalkControl::Skip if entry.is_dir() => {
100                    skipped.insert(entry.path().to_path_buf());
101                }
102                WalkControl::Continue | WalkControl::Skip => {}
103                WalkControl::Quit => stop.store(true, Ordering::Release),
104            }
105        }
106        shallow.tasks.retain(|task| !skipped.contains(&task.path));
107        for error in shallow.errors {
108            if cancellation.is_cancelled() || stop.load(Ordering::Acquire) {
109                break;
110            }
111            let control = visitor(WalkEvent::Error(&error));
112            let abort = self.options.error_policy == ErrorPolicy::Abort;
113            errors.push(error);
114            if abort || control == WalkControl::Quit {
115                stop.store(true, Ordering::Release);
116            }
117        }
118        if self.options.error_policy == ErrorPolicy::Abort && !errors.is_empty() {
119            return Err(errors.remove(0));
120        }
121        if shallow.tasks.is_empty() || cancellation.is_cancelled() || stop.load(Ordering::Acquire) {
122            return Ok(ParallelVisitReport {
123                visited,
124                errors,
125                quit: stop.load(Ordering::Acquire),
126                cancelled: cancellation.is_cancelled(),
127            });
128        }
129
130        let worker_count =
131            parallel_worker_count(self.parallelism, self.options.max_open, shallow.tasks.len());
132        let mut lanes = (0..worker_count)
133            .map(|_| Vec::<DirectoryTask>::new())
134            .collect::<Vec<_>>();
135        for (index, task) in shallow.tasks.into_iter().enumerate() {
136            lanes[index % worker_count].push(task);
137        }
138        let (sender, receiver) = mpsc::channel();
139        for (index, lane) in lanes.into_iter().enumerate() {
140            let sender = sender.clone();
141            let root = Arc::clone(&shallow.root);
142            let visitor = Arc::clone(&visitor);
143            let cancellation = cancellation.clone();
144            let stop = Arc::clone(&stop);
145            let options = self.options;
146            ThreadPool::global().execute(move || {
147                let report =
148                    visit_lane(lane, options, &root, &cancellation, &stop, visitor.as_ref());
149                let _ = sender.send((index, report));
150            });
151        }
152        drop(sender);
153        let mut completed = (0..worker_count).map(|_| None).collect::<Vec<_>>();
154        for (index, report) in receiver {
155            completed[index] = Some(report);
156        }
157        for report in completed {
158            let report = report.expect("every parallel visitor lane reports completion");
159            visited = visited.saturating_add(report.visited);
160            errors.extend(report.errors);
161        }
162        if self.options.error_policy == ErrorPolicy::Abort && !errors.is_empty() {
163            return Err(errors.remove(0));
164        }
165        Ok(ParallelVisitReport {
166            visited,
167            errors,
168            quit: stop.load(Ordering::Acquire),
169            cancelled: cancellation.is_cancelled(),
170        })
171    }
172}