use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ProgressEvent {
pub stage: &'static str,
pub current: u64,
pub total: u64,
}
pub type ProgressSink = Box<dyn Fn(ProgressEvent) + Send + Sync>;
static ARMED: AtomicBool = AtomicBool::new(false);
static SINK: Mutex<Option<ProgressSink>> = Mutex::new(None);
static VISION_TOTAL: AtomicU64 = AtomicU64::new(0);
static VISION_DONE: AtomicU64 = AtomicU64::new(0);
pub fn set_progress_sink(sink: Option<ProgressSink>) {
let armed = sink.is_some();
let mut slot = SINK.lock().unwrap_or_else(|e| e.into_inner());
*slot = sink;
ARMED.store(armed, Ordering::Relaxed);
}
#[must_use]
#[inline]
pub fn enabled() -> bool {
ARMED.load(Ordering::Relaxed)
}
#[inline]
pub fn emit(stage: &'static str, current: u64, total: u64) {
if !enabled() {
return;
}
emit_cold(ProgressEvent {
stage,
current,
total,
});
}
#[cold]
#[inline(never)]
fn emit_cold(event: ProgressEvent) {
let Ok(slot) = SINK.try_lock() else { return };
if let Some(sink) = slot.as_ref() {
sink(event);
}
}
pub fn vision_begin(total_blocks: u64) {
if !enabled() {
return;
}
VISION_TOTAL.store(total_blocks, Ordering::Relaxed);
VISION_DONE.store(0, Ordering::Relaxed);
emit("vision", 0, total_blocks);
}
#[inline]
pub fn vision_step() {
if !enabled() {
return;
}
let total = VISION_TOTAL.load(Ordering::Relaxed);
if total == 0 {
return;
}
let done = VISION_DONE.fetch_add(1, Ordering::Relaxed) + 1;
emit("vision", done.min(total), total);
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{Arc, Mutex as StdMutex};
static TEST_LOCK: Mutex<()> = Mutex::new(());
#[test]
fn disabled_by_default_and_free() {
let _guard = TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
set_progress_sink(None);
assert!(!enabled());
emit("decode", 7, 100);
vision_begin(12);
vision_step();
}
#[test]
fn events_reach_an_installed_sink() {
let _guard = TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let seen: Arc<StdMutex<Vec<ProgressEvent>>> = Arc::new(StdMutex::new(Vec::new()));
let sink_seen = Arc::clone(&seen);
set_progress_sink(Some(Box::new(move |ev| {
sink_seen.lock().unwrap_or_else(|e| e.into_inner()).push(ev);
})));
assert!(enabled());
emit("preprocess", 0, 0);
vision_begin(3);
vision_step();
vision_step();
vision_step();
vision_step();
emit("decode", 2, 256);
set_progress_sink(None);
assert!(!enabled());
emit("decode", 3, 256);
let events = seen.lock().unwrap_or_else(|e| e.into_inner()).clone();
let expected = [
ProgressEvent {
stage: "preprocess",
current: 0,
total: 0,
},
ProgressEvent {
stage: "vision",
current: 0,
total: 3,
},
ProgressEvent {
stage: "vision",
current: 1,
total: 3,
},
ProgressEvent {
stage: "vision",
current: 2,
total: 3,
},
ProgressEvent {
stage: "vision",
current: 3,
total: 3,
},
ProgressEvent {
stage: "vision",
current: 3,
total: 3,
},
ProgressEvent {
stage: "decode",
current: 2,
total: 256,
},
];
assert_eq!(events, expected);
}
#[test]
fn a_reentrant_sink_cannot_deadlock_the_forward() {
let _guard = TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let count = Arc::new(StdMutex::new(0usize));
let sink_count = Arc::clone(&count);
set_progress_sink(Some(Box::new(move |_| {
*sink_count.lock().unwrap_or_else(|e| e.into_inner()) += 1;
emit("decode", 0, 0);
})));
emit("decode", 1, 8);
set_progress_sink(None);
assert_eq!(*count.lock().unwrap_or_else(|e| e.into_inner()), 1);
}
}