qubit_progress/
auto_reporter.rs1use std::marker::PhantomData;
12use std::panic::AssertUnwindSafe;
13use std::panic::catch_unwind;
14use std::panic::resume_unwind;
15use std::sync::Arc;
16use std::sync::Weak;
17use std::sync::atomic::AtomicBool;
18use std::sync::atomic::Ordering;
19use std::sync::mpsc::Receiver;
20use std::sync::mpsc::SyncSender;
21use std::sync::mpsc::sync_channel;
22use std::thread;
23use std::thread::ScopedJoinHandle;
24
25use crate::AutoReporterError;
26use crate::EmissionError;
27use crate::Progress;
28use crate::WorkerPanic;
29
30#[must_use]
32pub struct AutoReporter<'scope, 'reporter> {
33 join: Option<ScopedJoinHandle<'scope, Result<(), EmissionError>>>,
35 inner: Option<Arc<AutoReporterInner>>,
37 status: AutoReporterStatus,
39 progress_borrow: PhantomData<&'scope mut Progress<'reporter>>,
41}
42
43impl<'scope, 'reporter> AutoReporter<'scope, 'reporter> {
44 #[must_use]
46 pub fn notifier(&self) -> ProgressNotifier {
47 ProgressNotifier {
48 inner: self.inner.as_ref().and_then(|inner| {
49 inner.notification_driven.then(|| Arc::downgrade(inner))
50 }),
51 }
52 }
53
54 #[must_use]
56 pub fn status(&self) -> AutoReporterStatus {
57 self.status.clone()
58 }
59
60 pub fn stop(mut self) -> Result<(), AutoReporterError> {
70 self.signal_stop();
71 match self.join_worker() {
72 Ok(result) => result.map_err(AutoReporterError::Emission),
73 Err(panic) => Err(AutoReporterError::Panicked(panic)),
74 }
75 }
76
77 fn signal_stop(&self) {
79 if let Some(inner) = &self.inner {
80 inner.stopped.store(true, Ordering::Release);
81 wake(&inner.wake_sender);
82 }
83 }
84
85 fn join_worker(
87 &mut self,
88 ) -> Result<Result<(), EmissionError>, WorkerPanic> {
89 let Some(join) = self.join.take() else {
90 return Ok(Ok(()));
91 };
92 join.join().map_err(WorkerPanic::new)
93 }
94}
95
96impl Drop for AutoReporter<'_, '_> {
97 fn drop(&mut self) {
99 self.signal_stop();
100 match self.join_worker() {
101 Ok(Ok(())) => {}
102 Ok(Err(_)) => self.status.mark_failed(),
103 Err(_) => self.status.mark_failed(),
104 }
105 }
106}
107
108#[derive(Clone)]
110pub struct ProgressNotifier {
111 inner: Option<Weak<AutoReporterInner>>,
113}
114
115impl ProgressNotifier {
116 pub fn notify(&self) {
122 let Some(inner) = self.inner.as_ref().and_then(Weak::upgrade) else {
123 return;
124 };
125 if inner.stopped.load(Ordering::Acquire) {
126 return;
127 }
128 inner.pending.store(true, Ordering::Release);
129 wake(&inner.wake_sender);
130 }
131}
132
133#[derive(Clone)]
135pub struct AutoReporterStatus {
136 failed: Arc<AtomicBool>,
138}
139
140impl AutoReporterStatus {
141 fn healthy() -> Self {
143 Self {
144 failed: Arc::new(AtomicBool::new(false)),
145 }
146 }
147
148 fn mark_failed(&self) {
150 self.failed.store(true, Ordering::Release);
151 }
152
153 #[must_use]
156 pub fn is_failed(&self) -> bool {
157 self.failed.load(Ordering::Acquire)
158 }
159}
160
161struct AutoReporterInner {
163 notification_driven: bool,
165 wake_sender: SyncSender<()>,
167 stopped: AtomicBool,
169 pending: AtomicBool,
171}
172
173pub(crate) fn spawn<'scope, 'env, 'reporter>(
175 progress: &'scope mut Progress<'reporter>,
176 scope: &'scope thread::Scope<'scope, 'env>,
177) -> AutoReporter<'scope, 'reporter>
178where
179 'reporter: 'scope,
180{
181 let status = AutoReporterStatus::healthy();
182 if !progress.is_enabled() {
183 return AutoReporter {
184 join: None,
185 inner: None,
186 status,
187 progress_borrow: PhantomData,
188 };
189 }
190 let (wake_sender, wake_receiver) = sync_channel(1);
191 let inner = Arc::new(AutoReporterInner {
192 notification_driven: progress.report_interval().is_zero(),
193 wake_sender,
194 stopped: AtomicBool::new(false),
195 pending: AtomicBool::new(false),
196 });
197 let worker_inner = Arc::clone(&inner);
198 let worker_status = status.clone();
199 let join = scope.spawn(move || {
200 match catch_unwind(AssertUnwindSafe(|| {
201 run(progress, Arc::clone(&worker_inner), wake_receiver)
202 })) {
203 Ok(result) => {
204 if result.is_err() {
205 worker_status.mark_failed();
206 worker_inner.stopped.store(true, Ordering::Release);
207 }
208 result
209 }
210 Err(payload) => {
211 worker_status.mark_failed();
212 worker_inner.stopped.store(true, Ordering::Release);
213 resume_unwind(payload)
214 }
215 }
216 });
217 AutoReporter {
218 join: Some(join),
219 inner: Some(inner),
220 status,
221 progress_borrow: PhantomData,
222 }
223}
224
225fn run(
227 progress: &mut Progress<'_>,
228 inner: Arc<AutoReporterInner>,
229 receiver: Receiver<()>,
230) -> Result<(), EmissionError> {
231 if progress.report_interval().is_zero() {
232 run_notified(progress, &inner, receiver)
233 } else {
234 run_heartbeat(progress, &inner, receiver)
235 }
236}
237
238fn run_notified(
240 progress: &mut Progress<'_>,
241 inner: &AutoReporterInner,
242 receiver: Receiver<()>,
243) -> Result<(), EmissionError> {
244 loop {
245 receiver
246 .recv()
247 .expect("notification sender must outlive the reporter worker");
248 if inner.pending.swap(false, Ordering::AcqRel) {
249 progress.report()?;
250 }
251 if inner.stopped.load(Ordering::Acquire) {
252 return Ok(());
253 }
254 }
255}
256
257fn run_heartbeat(
259 progress: &mut Progress<'_>,
260 inner: &AutoReporterInner,
261 receiver: Receiver<()>,
262) -> Result<(), EmissionError> {
263 loop {
264 if inner.stopped.load(Ordering::Acquire) {
265 return Ok(());
266 }
267 let timeout = progress.time_until_due();
268 if receiver.recv_timeout(timeout).is_ok() {
269 return Ok(());
270 }
271 progress.report_if_due()?;
272 }
273}
274
275fn wake(sender: &SyncSender<()>) {
277 let _ = sender.try_send(());
278}