Skip to main content

harn_test_runner/
scheduling.rs

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/// Execute already-discovered cases on owned worker threads.
47///
48/// The engine owns claiming, weighted resource permits, serial groups,
49/// fail-fast barriers, progress ordering, and worker lifecycle. The host
50/// adapter supplies only worker-local runtime construction and one-case
51/// execution, keeping CLI/package capability wiring out of this crate.
52#[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}