weavatrix_scan/parallel/
visit.rs1use 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#[derive(Debug)]
16pub enum WalkEvent<'a> {
17 Entry(&'a WalkEntry),
18 Error(&'a WalkError),
19}
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum WalkControl {
24 Continue,
25 Skip,
27 Quit,
29}
30
31#[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 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 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}