use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use parking_lot::Mutex;
pub type ProgressFn = Arc<dyn Fn(u64, u64) + Send + Sync>;
pub const REPORT_INTERVAL: f64 = 0.01;
#[derive(Debug, Clone, Default)]
pub struct CancelFlag(Arc<std::sync::atomic::AtomicBool>);
impl CancelFlag {
pub fn new() -> Self {
Self::default()
}
pub fn cancel(&self) {
self.0.store(true, Ordering::Relaxed);
}
pub fn is_cancelled(&self) -> bool {
self.0.load(Ordering::Relaxed)
}
}
pub struct ProgressTracker {
done: AtomicU64,
total: u64,
callback: Option<ProgressFn>,
last_reported: Mutex<f64>,
}
impl std::fmt::Debug for ProgressTracker {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ProgressTracker")
.field("done", &self.done())
.field("total", &self.total)
.field("reporting", &self.callback.is_some())
.finish()
}
}
impl ProgressTracker {
pub fn new(total: u64) -> Self {
Self::with_callback(total, None)
}
pub fn with_callback(total: u64, callback: Option<ProgressFn>) -> Self {
Self {
done: AtomicU64::new(0),
total,
callback,
last_reported: Mutex::new(0.0),
}
}
pub fn add(&self, value: u64) {
let current = self.done.fetch_add(value, Ordering::Relaxed) + value;
let Some(callback) = &self.callback else {
return;
};
let mut last = self.last_reported.lock();
let progress = if self.total > 0 {
current as f64 / self.total as f64
} else {
0.0
};
if progress >= *last + REPORT_INTERVAL {
*last = progress;
callback(current, self.total);
}
}
pub fn done_report(&self) {
let Some(callback) = &self.callback else {
return;
};
let mut last = self.last_reported.lock();
if *last < 1.0 {
*last = 1.0;
callback(self.total, self.total);
}
}
pub fn done(&self) -> u64 {
self.done.load(Ordering::Relaxed)
}
pub fn total(&self) -> u64 {
self.total
}
}
pub fn default_progress() -> ProgressFn {
Arc::new(|done, total| {
use std::io::{IsTerminal, Write};
let mut err = std::io::stderr();
if !err.is_terminal() {
return;
}
const WIDTH: usize = 40;
let fraction = if total == 0 {
1.0
} else {
(done as f64 / total as f64).clamp(0.0, 1.0)
};
let filled = (fraction * WIDTH as f64).round() as usize;
let _ = write!(
err,
"\r[{}{}] {:5.1}% {done}/{total} bp",
"#".repeat(filled),
"-".repeat(WIDTH - filled),
fraction * 100.0,
);
if done >= total {
let _ = writeln!(err);
}
let _ = err.flush();
})
}
#[cfg(test)]
mod tests {
use super::*;
type Recorder = (ProgressFn, Arc<Mutex<Vec<(u64, u64)>>>);
fn recording() -> Recorder {
let seen = Arc::new(Mutex::new(Vec::new()));
let sink = seen.clone();
(Arc::new(move |d, t| sink.lock().push((d, t))), seen)
}
#[test]
fn a_tracker_with_no_callback_just_counts() {
let tracker = ProgressTracker::new(100);
tracker.add(30);
tracker.add(20);
assert_eq!(tracker.done(), 50);
assert_eq!(tracker.total(), 100);
tracker.done_report(); }
#[test]
fn reports_only_once_the_interval_is_crossed() {
let (callback, seen) = recording();
let tracker = ProgressTracker::with_callback(1000, Some(callback));
for _ in 0..9 {
tracker.add(1);
}
assert!(seen.lock().is_empty());
tracker.add(1); assert_eq!(&*seen.lock(), &[(10, 1000)]);
}
#[test]
fn the_final_report_fires_once_and_says_the_total() {
let (callback, seen) = recording();
let tracker = ProgressTracker::with_callback(1000, Some(callback));
tracker.add(500);
tracker.done_report();
tracker.done_report();
assert_eq!(seen.lock().last(), Some(&(1000, 1000)));
assert_eq!(seen.lock().iter().filter(|(d, _)| *d == 1000).count(), 1);
}
#[test]
fn a_request_that_finished_exactly_does_not_report_twice() {
let (callback, seen) = recording();
let tracker = ProgressTracker::with_callback(100, Some(callback));
tracker.add(100);
assert_eq!(&*seen.lock(), &[(100, 100)]);
tracker.done_report();
assert_eq!(seen.lock().len(), 1);
}
#[test]
fn a_zero_total_never_divides_and_still_completes() {
let (callback, seen) = recording();
let tracker = ProgressTracker::with_callback(0, Some(callback));
tracker.add(5);
assert!(seen.lock().is_empty());
tracker.done_report();
assert_eq!(&*seen.lock(), &[(0, 0)]);
}
#[test]
fn concurrent_adds_are_not_lost_and_reports_are_serialised() {
let (callback, seen) = recording();
let tracker = Arc::new(ProgressTracker::with_callback(8000, Some(callback)));
let threads: Vec<_> = (0..8)
.map(|_| {
let tracker = tracker.clone();
std::thread::spawn(move || {
for _ in 0..1000 {
tracker.add(1);
}
})
})
.collect();
for t in threads {
t.join().unwrap();
}
assert_eq!(tracker.done(), 8000);
let seen = seen.lock();
assert!(!seen.is_empty());
assert!(seen.len() <= 100, "{} reports", seen.len());
assert!(seen.windows(2).all(|w| w[0].0 <= w[1].0));
}
}