use std::marker::PhantomData;
use std::panic::AssertUnwindSafe;
use std::panic::catch_unwind;
use std::panic::resume_unwind;
use std::sync::Arc;
use std::sync::Weak;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering;
use std::sync::mpsc::Receiver;
use std::sync::mpsc::SyncSender;
use std::sync::mpsc::sync_channel;
use std::thread;
use std::thread::ScopedJoinHandle;
use crate::AutoReporterError;
use crate::EmissionError;
use crate::Progress;
use crate::WorkerPanic;
#[must_use]
pub struct AutoReporter<'scope, 'reporter> {
join: Option<ScopedJoinHandle<'scope, Result<(), EmissionError>>>,
inner: Option<Arc<AutoReporterInner>>,
status: AutoReporterStatus,
progress_borrow: PhantomData<&'scope mut Progress<'reporter>>,
}
impl<'scope, 'reporter> AutoReporter<'scope, 'reporter> {
#[must_use]
pub fn notifier(&self) -> ProgressNotifier {
ProgressNotifier {
inner: self.inner.as_ref().and_then(|inner| {
inner.notification_driven.then(|| Arc::downgrade(inner))
}),
}
}
#[must_use]
pub fn status(&self) -> AutoReporterStatus {
self.status.clone()
}
pub fn stop(mut self) -> Result<(), AutoReporterError> {
self.signal_stop();
match self.join_worker() {
Ok(result) => result.map_err(AutoReporterError::Emission),
Err(panic) => Err(AutoReporterError::Panicked(panic)),
}
}
fn signal_stop(&self) {
if let Some(inner) = &self.inner {
inner.stopped.store(true, Ordering::Release);
wake(&inner.wake_sender);
}
}
fn join_worker(
&mut self,
) -> Result<Result<(), EmissionError>, WorkerPanic> {
let Some(join) = self.join.take() else {
return Ok(Ok(()));
};
join.join().map_err(WorkerPanic::new)
}
}
impl Drop for AutoReporter<'_, '_> {
fn drop(&mut self) {
self.signal_stop();
match self.join_worker() {
Ok(Ok(())) => {}
Ok(Err(_)) => self.status.mark_failed(),
Err(_) => self.status.mark_failed(),
}
}
}
#[derive(Clone)]
pub struct ProgressNotifier {
inner: Option<Weak<AutoReporterInner>>,
}
impl ProgressNotifier {
pub fn notify(&self) {
let Some(inner) = self.inner.as_ref().and_then(Weak::upgrade) else {
return;
};
if inner.stopped.load(Ordering::Acquire) {
return;
}
inner.pending.store(true, Ordering::Release);
wake(&inner.wake_sender);
}
}
#[derive(Clone)]
pub struct AutoReporterStatus {
failed: Arc<AtomicBool>,
}
impl AutoReporterStatus {
fn healthy() -> Self {
Self {
failed: Arc::new(AtomicBool::new(false)),
}
}
fn mark_failed(&self) {
self.failed.store(true, Ordering::Release);
}
#[must_use]
pub fn is_failed(&self) -> bool {
self.failed.load(Ordering::Acquire)
}
}
struct AutoReporterInner {
notification_driven: bool,
wake_sender: SyncSender<()>,
stopped: AtomicBool,
pending: AtomicBool,
}
pub(crate) fn spawn<'scope, 'env, 'reporter>(
progress: &'scope mut Progress<'reporter>,
scope: &'scope thread::Scope<'scope, 'env>,
) -> AutoReporter<'scope, 'reporter>
where
'reporter: 'scope,
{
let status = AutoReporterStatus::healthy();
if !progress.is_enabled() {
return AutoReporter {
join: None,
inner: None,
status,
progress_borrow: PhantomData,
};
}
let (wake_sender, wake_receiver) = sync_channel(1);
let inner = Arc::new(AutoReporterInner {
notification_driven: progress.report_interval().is_zero(),
wake_sender,
stopped: AtomicBool::new(false),
pending: AtomicBool::new(false),
});
let worker_inner = Arc::clone(&inner);
let worker_status = status.clone();
let join = scope.spawn(move || {
match catch_unwind(AssertUnwindSafe(|| {
run(progress, Arc::clone(&worker_inner), wake_receiver)
})) {
Ok(result) => {
if result.is_err() {
worker_status.mark_failed();
worker_inner.stopped.store(true, Ordering::Release);
}
result
}
Err(payload) => {
worker_status.mark_failed();
worker_inner.stopped.store(true, Ordering::Release);
resume_unwind(payload)
}
}
});
AutoReporter {
join: Some(join),
inner: Some(inner),
status,
progress_borrow: PhantomData,
}
}
fn run(
progress: &mut Progress<'_>,
inner: Arc<AutoReporterInner>,
receiver: Receiver<()>,
) -> Result<(), EmissionError> {
if progress.report_interval().is_zero() {
run_notified(progress, &inner, receiver)
} else {
run_heartbeat(progress, &inner, receiver)
}
}
fn run_notified(
progress: &mut Progress<'_>,
inner: &AutoReporterInner,
receiver: Receiver<()>,
) -> Result<(), EmissionError> {
loop {
receiver
.recv()
.expect("notification sender must outlive the reporter worker");
if inner.pending.swap(false, Ordering::AcqRel) {
progress.report()?;
}
if inner.stopped.load(Ordering::Acquire) {
return Ok(());
}
}
}
fn run_heartbeat(
progress: &mut Progress<'_>,
inner: &AutoReporterInner,
receiver: Receiver<()>,
) -> Result<(), EmissionError> {
loop {
if inner.stopped.load(Ordering::Acquire) {
return Ok(());
}
let timeout = progress.time_until_due();
if receiver.recv_timeout(timeout).is_ok() {
return Ok(());
}
progress.report_if_due()?;
}
}
fn wake(sender: &SyncSender<()>) {
let _ = sender.try_send(());
}