Skip to main content

ruccl/rank/
work.rs

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/// Asynchronous collective completion handle. Dropping it detaches the host
42/// worker but does not cancel or reorder an already submitted collective.
43#[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}