1use std::collections::HashSet;
2use std::sync::atomic::{AtomicBool, Ordering};
3use std::sync::{Arc, Condvar, Mutex};
4use std::thread;
5
6use crate::{TestCase, TestResult};
7
8#[derive(Clone, Debug)]
9pub enum TestRunEvent {
10 SuiteDiscovered {
11 total_tests: usize,
12 total_files: usize,
13 parallel: bool,
14 workers: usize,
15 },
16 LargeSequentialSuite {
17 total_tests: usize,
18 total_files: usize,
19 },
20 TestStarted {
21 name: String,
22 file: String,
23 test_index: usize,
24 total_tests: usize,
25 },
26 TestFinished(TestResult),
27}
28
29pub type TestRunProgress = Arc<dyn Fn(TestRunEvent) + Send + Sync>;
30
31#[doc(hidden)]
32pub struct ParallelCaseResults {
33 pub cases: Vec<TestResult>,
34 pub infrastructure_errors: Vec<TestResult>,
35}
36
37#[doc(hidden)]
38pub struct ParallelRunOptions {
39 pub workers: usize,
40 pub total_tests: usize,
41 pub stack_size: usize,
42 pub fail_fast: bool,
43 pub progress: Option<TestRunProgress>,
44}
45
46#[doc(hidden)]
53pub fn execute_parallel_cases<W, Init, Execute>(
54 cases: Vec<TestCase>,
55 options: ParallelRunOptions,
56 init_worker: Init,
57 execute: Execute,
58) -> ParallelCaseResults
59where
60 W: Send + 'static,
61 Init: Fn(usize) -> Result<W, String> + Send + Sync + 'static,
62 Execute: Fn(&mut W, &TestCase) -> TestResult + Send + Sync + 'static,
63{
64 let ParallelRunOptions {
65 workers,
66 total_tests,
67 stack_size,
68 fail_fast,
69 progress,
70 } = options;
71 let queue = Arc::new(Mutex::new(cases));
72 let gate = Arc::new(ResourceGate::new(workers));
73 let results: Arc<Mutex<Vec<TestResult>>> = Arc::new(Mutex::new(Vec::new()));
74 let infrastructure_errors: Arc<Mutex<Vec<TestResult>>> = Arc::new(Mutex::new(Vec::new()));
75 let cancelled = Arc::new(AtomicBool::new(false));
76 let completed = Arc::new(Mutex::new(0usize));
77 let init_worker = Arc::new(init_worker);
78 let execute = Arc::new(execute);
79
80 let mut handles = Vec::with_capacity(workers);
81 for worker_idx in 0..workers {
82 let queue = Arc::clone(&queue);
83 let gate = Arc::clone(&gate);
84 let results = Arc::clone(&results);
85 let infrastructure_errors = Arc::clone(&infrastructure_errors);
86 let cancelled = Arc::clone(&cancelled);
87 let completed = Arc::clone(&completed);
88 let progress = progress.clone();
89 let init_worker = Arc::clone(&init_worker);
90 let execute = Arc::clone(&execute);
91 let handle = thread::Builder::new()
92 .name(format!("harn-test-worker-{worker_idx}"))
93 .stack_size(stack_size)
94 .spawn(move || {
95 let mut worker = match init_worker(worker_idx) {
96 Ok(worker) => worker,
97 Err(error) => {
98 infrastructure_errors.lock().unwrap().push(TestResult {
99 name: "<worker error>".to_string(),
100 file: String::new(),
101 passed: false,
102 error: Some(error),
103 captured_output: None,
104 timeout: None,
105 duration_ms: 0,
106 phases: None,
107 });
108 return;
109 }
110 };
111 loop {
112 let Some(case) = claim_next_case(&queue, &cancelled, fail_fast) else {
113 break;
114 };
115 let _guard = gate.acquire(case.weight, case.serial_group.as_deref());
116 if fail_fast && cancelled.load(Ordering::Acquire) {
117 break;
118 }
119 let test_index = next_test_index(&completed);
120 emit_progress(
121 &progress,
122 TestRunEvent::TestStarted {
123 name: case.name.clone(),
124 file: case.file.display().to_string(),
125 test_index,
126 total_tests,
127 },
128 );
129 let result = execute(&mut worker, &case);
130 if fail_fast && !result.passed {
131 cancelled.store(true, Ordering::Release);
132 }
133 emit_progress(&progress, TestRunEvent::TestFinished(result.clone()));
134 results.lock().unwrap().push(result);
135 }
136 })
137 .expect("spawning a harn-test worker thread should succeed");
138 handles.push(handle);
139 }
140 for handle in handles {
141 let _ = handle.join();
142 }
143
144 ParallelCaseResults {
145 cases: unwrap_or_clone(results),
146 infrastructure_errors: unwrap_or_clone(infrastructure_errors),
147 }
148}
149
150fn unwrap_or_clone(values: Arc<Mutex<Vec<TestResult>>>) -> Vec<TestResult> {
151 Arc::try_unwrap(values)
152 .map(|mutex| mutex.into_inner().unwrap_or_default())
153 .unwrap_or_else(|arc| arc.lock().unwrap().clone())
154}
155
156fn emit_progress(progress: &Option<TestRunProgress>, event: TestRunEvent) {
157 if let Some(callback) = progress {
158 callback(event);
159 }
160}
161
162fn claim_next_case(
163 queue: &Mutex<Vec<TestCase>>,
164 cancelled: &AtomicBool,
165 fail_fast: bool,
166) -> Option<TestCase> {
167 let mut queue = queue.lock().unwrap();
168 if fail_fast && cancelled.load(Ordering::Acquire) {
169 None
170 } else {
171 queue.pop()
172 }
173}
174
175fn next_test_index(counter: &Mutex<usize>) -> usize {
176 let mut guard = counter.lock().unwrap();
177 *guard += 1;
178 *guard
179}
180
181struct ResourceGate {
182 state: Mutex<GateState>,
183 cond: Condvar,
184 capacity: usize,
185}
186
187struct GateState {
188 available: usize,
189 busy_groups: HashSet<String>,
190}
191
192struct GateGuard<'a> {
193 gate: &'a ResourceGate,
194 weight: usize,
195 group: Option<String>,
196}
197
198impl ResourceGate {
199 fn new(capacity: usize) -> Self {
200 Self {
201 state: Mutex::new(GateState {
202 available: capacity,
203 busy_groups: HashSet::new(),
204 }),
205 cond: Condvar::new(),
206 capacity,
207 }
208 }
209
210 fn acquire(&self, weight: usize, group: Option<&str>) -> GateGuard<'_> {
211 let weight = weight.min(self.capacity).max(1);
212 let mut state = self.state.lock().unwrap();
213 loop {
214 let group_free = group.is_none_or(|name| !state.busy_groups.contains(name));
215 if state.available >= weight && group_free {
216 state.available -= weight;
217 if let Some(name) = group {
218 state.busy_groups.insert(name.to_string());
219 }
220 return GateGuard {
221 gate: self,
222 weight,
223 group: group.map(str::to_owned),
224 };
225 }
226 state = self.cond.wait(state).unwrap();
227 }
228 }
229
230 #[cfg(test)]
231 fn try_acquire(&self, weight: usize, group: Option<&str>) -> Option<GateGuard<'_>> {
232 let weight = weight.min(self.capacity).max(1);
233 let mut state = self.state.lock().unwrap();
234 let group_free = group.is_none_or(|name| !state.busy_groups.contains(name));
235 if state.available < weight || !group_free {
236 return None;
237 }
238 state.available -= weight;
239 if let Some(name) = group {
240 state.busy_groups.insert(name.to_string());
241 }
242 Some(GateGuard {
243 gate: self,
244 weight,
245 group: group.map(str::to_owned),
246 })
247 }
248}
249
250impl Drop for GateGuard<'_> {
251 fn drop(&mut self) {
252 let mut state = self.gate.state.lock().unwrap();
253 state.available += self.weight;
254 if let Some(group) = self.group.as_deref() {
255 state.busy_groups.remove(group);
256 }
257 self.gate.cond.notify_all();
258 }
259}
260
261#[cfg(test)]
262mod tests {
263 use std::path::Path;
264 use std::sync::{Arc, Mutex};
265
266 use super::{execute_parallel_cases, ParallelRunOptions, ResourceGate, TestRunEvent};
267 use crate::{extract_cases_from_program, parse_program, TestResult};
268
269 #[test]
270 fn serial_groups_and_weights_are_structural_not_timing_based() {
271 let gate = ResourceGate::new(2);
272 let login = gate.acquire(1, Some("login"));
273 assert!(gate.try_acquire(1, Some("login")).is_none());
274 assert!(gate.try_acquire(1, Some("independent")).is_some());
275 drop(login);
276 assert!(gate.try_acquire(1, Some("login")).is_some());
277
278 let gate = ResourceGate::new(2);
279 let all = gate.acquire(99, None);
280 assert!(gate.try_acquire(1, None).is_none());
281 drop(all);
282 assert!(gate.try_acquire(1, None).is_some());
283 }
284
285 #[test]
286 fn fail_fast_stops_claiming_and_progress_is_event_driven() {
287 let source = Arc::new(
288 "pipeline test_one(task) { assert(true) }\n\
289 pipeline test_two(task) { assert(true) }\n"
290 .to_string(),
291 );
292 let program = Arc::new(parse_program(&source).unwrap());
293 let cases = extract_cases_from_program(
294 Path::new("test_scheduler.harn"),
295 &source,
296 &program,
297 None,
298 usize::MAX,
299 )
300 .unwrap();
301 let events = Arc::new(Mutex::new(Vec::new()));
302 let captured = Arc::clone(&events);
303 let progress = Arc::new(move |event| {
304 captured.lock().unwrap().push(match event {
305 TestRunEvent::TestStarted { .. } => "started",
306 TestRunEvent::TestFinished(_) => "finished",
307 _ => "suite",
308 });
309 });
310
311 let run = execute_parallel_cases(
312 cases,
313 ParallelRunOptions {
314 workers: 1,
315 total_tests: 2,
316 stack_size: 2 * 1024 * 1024,
317 fail_fast: true,
318 progress: Some(progress),
319 },
320 |_| Ok(()),
321 |_, case| TestResult {
322 name: case.name.clone(),
323 file: case.file.display().to_string(),
324 passed: false,
325 error: Some("deterministic failure".to_string()),
326 captured_output: None,
327 timeout: None,
328 duration_ms: 0,
329 phases: None,
330 },
331 );
332
333 assert_eq!(run.cases.len(), 1);
334 assert!(run.infrastructure_errors.is_empty());
335 assert_eq!(*events.lock().unwrap(), ["started", "finished"]);
336 }
337}