1use std::{
12 marker::PhantomData,
13 panic::{
14 AssertUnwindSafe,
15 catch_unwind,
16 resume_unwind,
17 },
18 sync::{
19 Arc,
20 Weak,
21 atomic::{
22 AtomicBool,
23 Ordering,
24 },
25 mpsc::{
26 Receiver,
27 SyncSender,
28 sync_channel,
29 },
30 },
31 thread::{
32 self,
33 ScopedJoinHandle,
34 },
35};
36
37use crate::{
38 AutoReporterError,
39 EmissionError,
40 Progress,
41 WorkerPanic,
42};
43
44#[must_use]
46pub struct AutoReporter<'scope, 'reporter> {
47 join: Option<ScopedJoinHandle<'scope, Result<(), EmissionError>>>,
49 inner: Option<Arc<AutoReporterInner>>,
51 status: AutoReporterStatus,
53 progress_borrow: PhantomData<&'scope mut Progress<'reporter>>,
55}
56
57impl<'scope, 'reporter> AutoReporter<'scope, 'reporter> {
58 #[must_use]
60 pub fn notifier(&self) -> ProgressNotifier {
61 ProgressNotifier {
62 inner: self.inner.as_ref().and_then(|inner| {
63 inner.notification_driven.then(|| Arc::downgrade(inner))
64 }),
65 }
66 }
67
68 #[must_use]
70 pub fn status(&self) -> AutoReporterStatus {
71 self.status.clone()
72 }
73
74 pub fn stop(mut self) -> Result<(), AutoReporterError> {
84 self.signal_stop();
85 match self.join_worker() {
86 Ok(result) => result.map_err(AutoReporterError::Emission),
87 Err(panic) => Err(AutoReporterError::Panicked(panic)),
88 }
89 }
90
91 fn signal_stop(&self) {
93 if let Some(inner) = &self.inner {
94 inner.stopped.store(true, Ordering::Release);
95 wake(&inner.wake_sender);
96 }
97 }
98
99 fn join_worker(
101 &mut self,
102 ) -> Result<Result<(), EmissionError>, WorkerPanic> {
103 let Some(join) = self.join.take() else {
104 return Ok(Ok(()));
105 };
106 join.join().map_err(WorkerPanic::new)
107 }
108}
109
110impl Drop for AutoReporter<'_, '_> {
111 fn drop(&mut self) {
113 self.signal_stop();
114 match self.join_worker() {
115 Ok(Ok(())) => {}
116 Ok(Err(_)) => self.status.mark_failed(),
117 Err(_) => self.status.mark_failed(),
118 }
119 }
120}
121
122#[derive(Clone)]
124pub struct ProgressNotifier {
125 inner: Option<Weak<AutoReporterInner>>,
127}
128
129impl ProgressNotifier {
130 pub fn notify(&self) {
136 let Some(inner) = self.inner.as_ref().and_then(Weak::upgrade) else {
137 return;
138 };
139 if inner.stopped.load(Ordering::Acquire) {
140 return;
141 }
142 inner.pending.store(true, Ordering::Release);
143 wake(&inner.wake_sender);
144 }
145}
146
147#[derive(Clone)]
149pub struct AutoReporterStatus {
150 failed: Arc<AtomicBool>,
152}
153
154impl AutoReporterStatus {
155 fn healthy() -> Self {
157 Self {
158 failed: Arc::new(AtomicBool::new(false)),
159 }
160 }
161
162 fn mark_failed(&self) {
164 self.failed.store(true, Ordering::Release);
165 }
166
167 #[must_use]
170 pub fn is_failed(&self) -> bool {
171 self.failed.load(Ordering::Acquire)
172 }
173}
174
175struct AutoReporterInner {
177 notification_driven: bool,
179 wake_sender: SyncSender<()>,
181 stopped: AtomicBool,
183 pending: AtomicBool,
185}
186
187pub(crate) fn spawn<'scope, 'env, 'reporter>(
189 progress: &'scope mut Progress<'reporter>,
190 scope: &'scope thread::Scope<'scope, 'env>,
191) -> AutoReporter<'scope, 'reporter>
192where
193 'reporter: 'scope,
194{
195 let status = AutoReporterStatus::healthy();
196 if !progress.is_enabled() {
197 return AutoReporter {
198 join: None,
199 inner: None,
200 status,
201 progress_borrow: PhantomData,
202 };
203 }
204 let (wake_sender, wake_receiver) = sync_channel(1);
205 let inner = Arc::new(AutoReporterInner {
206 notification_driven: progress.report_interval().is_zero(),
207 wake_sender,
208 stopped: AtomicBool::new(false),
209 pending: AtomicBool::new(false),
210 });
211 let worker_inner = Arc::clone(&inner);
212 let worker_status = status.clone();
213 let join = scope.spawn(move || {
214 match catch_unwind(AssertUnwindSafe(|| {
215 run(progress, Arc::clone(&worker_inner), wake_receiver)
216 })) {
217 Ok(result) => {
218 if result.is_err() {
219 worker_status.mark_failed();
220 worker_inner.stopped.store(true, Ordering::Release);
221 }
222 result
223 }
224 Err(payload) => {
225 worker_status.mark_failed();
226 worker_inner.stopped.store(true, Ordering::Release);
227 resume_unwind(payload)
228 }
229 }
230 });
231 AutoReporter {
232 join: Some(join),
233 inner: Some(inner),
234 status,
235 progress_borrow: PhantomData,
236 }
237}
238
239fn run(
241 progress: &mut Progress<'_>,
242 inner: Arc<AutoReporterInner>,
243 receiver: Receiver<()>,
244) -> Result<(), EmissionError> {
245 if progress.report_interval().is_zero() {
246 run_notified(progress, &inner, receiver)
247 } else {
248 run_heartbeat(progress, &inner, receiver)
249 }
250}
251
252fn run_notified(
254 progress: &mut Progress<'_>,
255 inner: &AutoReporterInner,
256 receiver: Receiver<()>,
257) -> Result<(), EmissionError> {
258 loop {
259 receiver
260 .recv()
261 .expect("notification sender must outlive the reporter worker");
262 if inner.pending.swap(false, Ordering::AcqRel) {
263 progress.report()?;
264 }
265 if inner.stopped.load(Ordering::Acquire) {
266 return Ok(());
267 }
268 }
269}
270
271fn run_heartbeat(
273 progress: &mut Progress<'_>,
274 inner: &AutoReporterInner,
275 receiver: Receiver<()>,
276) -> Result<(), EmissionError> {
277 loop {
278 if inner.stopped.load(Ordering::Acquire) {
279 return Ok(());
280 }
281 let timeout = progress.time_until_due();
282 if receiver.recv_timeout(timeout).is_ok() {
283 return Ok(());
284 }
285 progress.report_if_due()?;
286 }
287}
288
289fn wake(sender: &SyncSender<()>) {
291 let _ = sender.try_send(());
292}