use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Default, Debug, Clone)]
pub struct ScopeTracker {
inner: Arc<AtomicUsize>,
}
impl ScopeTracker {
pub fn new() -> Self {
Self {
inner: Arc::new(AtomicUsize::new(0)),
}
}
pub fn measure_scope(&self) -> ScopeTrackerGuard {
ScopeTrackerGuard::measure(self)
}
pub fn get(&self, ordering: Ordering) -> usize {
self.inner.load(ordering)
}
}
const COUNT_SIZE: usize = 1;
#[must_use = "dropping this guard immediately decrements the scope counter"]
pub struct ScopeTrackerGuard {
scope_tracker: ScopeTracker,
}
impl ScopeTrackerGuard {
fn measure(scope_tracker: &ScopeTracker) -> Self {
let scope_tracker = scope_tracker.clone();
scope_tracker.inner.fetch_add(COUNT_SIZE, Ordering::SeqCst);
Self { scope_tracker }
}
}
impl Drop for ScopeTrackerGuard {
fn drop(&mut self) {
self.scope_tracker
.inner
.fetch_sub(COUNT_SIZE, Ordering::SeqCst);
}
}
#[cfg(test)]
mod test {
use std::sync::atomic::{AtomicBool, Ordering};
use std::thread::{self, JoinHandle};
use std::time::Duration;
use super::*;
#[test]
fn test_scope_tracker() {
let counter = ScopeTracker::new();
{
let _measure_guard = counter.measure_scope();
assert_eq!(counter.get(Ordering::SeqCst), 1);
}
assert_eq!(counter.get(Ordering::SeqCst), 0);
}
#[test]
fn test_scope_tracker_loop() {
let counter = ScopeTracker::new();
for _ in 0..100 {
let _measure_guard = counter.measure_scope();
assert_eq!(counter.get(Ordering::SeqCst), 1);
}
assert_eq!(counter.get(Ordering::SeqCst), 0);
}
#[test]
fn test_scope_tracker_threads() {
let counter = ScopeTracker::new();
let run = Arc::new(AtomicBool::new(true));
let mut handles: Vec<JoinHandle<()>> = vec![];
let started_threads = Arc::new(AtomicUsize::new(0));
const LEN: usize = 20;
for _ in 0..LEN {
let counter_clone = counter.clone();
let run_clone = run.clone();
let started_threads_clone = started_threads.clone();
let handle = thread::spawn(move || {
let _guard = counter_clone.measure_scope();
started_threads_clone.fetch_add(1, Ordering::Relaxed);
while run_clone.load(Ordering::Relaxed) {
thread::sleep(Duration::from_secs(1));
}
});
handles.push(handle);
}
while started_threads.load(Ordering::Relaxed) < LEN {
thread::sleep(Duration::from_secs(1));
}
assert_eq!(counter.get(Ordering::SeqCst), LEN);
run.store(false, Ordering::Release);
for handle in handles {
handle.join().unwrap();
}
}
}