1use std::error::Error;
2use std::fmt::{Display, Formatter};
3use std::panic::{AssertUnwindSafe, catch_unwind};
4use std::sync::mpsc;
5use std::sync::{Arc, Condvar, Mutex};
6use std::thread::{self, JoinHandle};
7use std::time::Duration;
8
9#[cfg(test)]
10mod tests;
11
12#[derive(Debug)]
13pub enum WorkError {
14 WorkerStart(String),
15 WorkerQueueClosed,
16 WorkerPanic,
17 Poisoned,
18}
19
20impl Display for WorkError {
21 fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
22 match self {
23 Self::WorkerStart(error) => {
24 write!(formatter, "cannot start collective worker: {error}")
25 }
26 Self::WorkerQueueClosed => write!(formatter, "collective worker queue is closed"),
27 Self::WorkerPanic => write!(formatter, "collective worker panicked"),
28 Self::Poisoned => write!(formatter, "collective synchronization state is poisoned"),
29 }
30 }
31}
32
33impl Error for WorkError {}
34
35#[derive(Debug)]
36struct WorkCompletion<T, E> {
37 result: Mutex<Option<Result<T, E>>>,
38 changed: Condvar,
39}
40
41#[derive(Debug)]
44pub struct CollectiveWork<T, E = WorkError> {
45 completion: Arc<WorkCompletion<T, E>>,
46 worker: Option<JoinHandle<()>>,
47}
48
49type OrderedTask = Box<dyn FnOnce() + Send + 'static>;
50
51#[derive(Clone)]
52pub struct OrderedWorkQueue {
53 sender: mpsc::Sender<OrderedTask>,
54}
55
56impl std::fmt::Debug for OrderedWorkQueue {
57 fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
58 formatter.write_str("OrderedWorkQueue")
59 }
60}
61
62impl OrderedWorkQueue {
63 pub fn new(name: &'static str) -> Result<Self, WorkError> {
64 let (sender, receiver) = mpsc::channel::<OrderedTask>();
65 thread::Builder::new()
66 .name(name.into())
67 .spawn(move || {
68 while let Ok(task) = receiver.recv() {
69 task();
70 }
71 })
72 .map_err(|error| WorkError::WorkerStart(error.to_string()))?;
73 Ok(Self { sender })
74 }
75
76 pub fn submit<T: Send + 'static, E: From<WorkError> + Send + 'static>(
77 &self,
78 task: impl FnOnce() -> Result<T, E> + Send + 'static,
79 ) -> Result<CollectiveWork<T, E>, E> {
80 let completion = Arc::new(WorkCompletion {
81 result: Mutex::new(None),
82 changed: Condvar::new(),
83 });
84 let target = Arc::clone(&completion);
85 self.sender
86 .send(Box::new(move || {
87 let result = catch_unwind(AssertUnwindSafe(task))
88 .map_err(|_| E::from(WorkError::WorkerPanic))
89 .and_then(|result| result);
90 if let Ok(mut state) = target.result.lock() {
91 *state = Some(result);
92 target.changed.notify_all();
93 }
94 }))
95 .map_err(|_| E::from(WorkError::WorkerQueueClosed))?;
96 Ok(CollectiveWork {
97 completion,
98 worker: None,
99 })
100 }
101}
102
103impl<T: Send + 'static, E: From<WorkError> + Send + 'static> CollectiveWork<T, E> {
104 pub fn spawn(
105 name: &'static str,
106 task: impl FnOnce() -> Result<T, E> + Send + 'static,
107 ) -> Result<Self, E> {
108 let completion = Arc::new(WorkCompletion {
109 result: Mutex::new(None),
110 changed: Condvar::new(),
111 });
112 let target = Arc::clone(&completion);
113 let worker = thread::Builder::new()
114 .name(name.into())
115 .spawn(move || {
116 let result = catch_unwind(AssertUnwindSafe(task))
117 .map_err(|_| E::from(WorkError::WorkerPanic))
118 .and_then(|result| result);
119 if let Ok(mut state) = target.result.lock() {
120 *state = Some(result);
121 target.changed.notify_all();
122 }
123 })
124 .map_err(|error| E::from(WorkError::WorkerStart(error.to_string())))?;
125 Ok(Self {
126 completion,
127 worker: Some(worker),
128 })
129 }
130
131 pub fn is_complete(&self) -> bool {
132 self.completion
133 .result
134 .lock()
135 .map(|result| result.is_some())
136 .unwrap_or(true)
137 }
138
139 pub fn wait_for(&self, timeout: Duration) -> Result<bool, E> {
140 let result = self
141 .completion
142 .result
143 .lock()
144 .map_err(|_| E::from(WorkError::Poisoned))?;
145 let (result, _) = self
146 .completion
147 .changed
148 .wait_timeout_while(result, timeout, |result| result.is_none())
149 .map_err(|_| E::from(WorkError::Poisoned))?;
150 Ok(result.is_some())
151 }
152
153 pub fn wait(mut self) -> Result<T, E> {
154 let result = {
155 let mut state = self
156 .completion
157 .result
158 .lock()
159 .map_err(|_| E::from(WorkError::Poisoned))?;
160 while state.is_none() {
161 state = self
162 .completion
163 .changed
164 .wait(state)
165 .map_err(|_| E::from(WorkError::Poisoned))?;
166 }
167 state.take().expect("collective wait exits only when ready")
168 };
169 if let Some(worker) = self.worker.take() {
170 worker.join().map_err(|_| E::from(WorkError::WorkerPanic))?;
171 }
172 result
173 }
174}