1use parking_lot::Mutex;
2use std::sync::{
3 atomic::{AtomicBool, Ordering},
4 Arc,
5};
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
8#[serde(transparent)]
9pub struct WorkerId(u64);
10impl WorkerId {
11 pub const fn new(value: u64) -> Self {
12 Self(value)
13 }
14 pub const fn get(self) -> u64 {
15 self.0
16 }
17}
18#[derive(Debug, Default, Clone, serde::Serialize, serde::Deserialize)]
19#[serde(transparent)]
20pub struct WorkerIds(u64);
21impl WorkerIds {
22 pub fn allocate(&mut self) -> Result<WorkerId, Failure> {
23 let next = self.0.checked_add(1).ok_or_else(|| {
24 Failure::new(
25 FailureKind::IdentityExhausted,
26 "worker request IDs exhausted",
27 )
28 })?;
29 self.0 = next;
30 Ok(WorkerId::new(next))
31 }
32}
33#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
34pub enum FailureKind {
35 IdentityExhausted,
36 ThreadStart,
37 Panic,
38 Io,
39 Spawn,
40 Wait,
41 Exit,
42 InvalidInput,
43 Unavailable,
44 Protocol,
45 Disconnected,
46}
47#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
48pub struct Failure {
49 pub kind: FailureKind,
50 pub message: String,
51}
52impl Failure {
53 pub fn new(kind: FailureKind, message: impl Into<String>) -> Self {
54 Self {
55 kind,
56 message: message.into(),
57 }
58 }
59}
60#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
61pub enum CancelReason {
62 Superseded,
63 OwnerClosed,
64 Dismissed,
65 Shutdown,
66}
67#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
68pub enum Outcome<T> {
69 Success(T),
70 Failed {
71 failure: Failure,
72 partial: Option<T>,
73 },
74 Cancelled(CancelReason),
75}
76impl<T> Outcome<T> {
77 pub fn failed(kind: FailureKind, message: impl Into<String>) -> Self {
78 Self::Failed {
79 failure: Failure::new(kind, message),
80 partial: None,
81 }
82 }
83}
84#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
85pub struct Ticket<K> {
86 pub request: WorkerId,
87 pub key: K,
88}
89#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
90pub struct Completion<K, T> {
91 pub ticket: Ticket<K>,
92 pub outcome: Outcome<T>,
93}
94#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
95pub enum Load<K> {
96 Idle,
97 Running(Ticket<K>),
98 Ready(K),
99 Failed { key: K, failure: Failure },
100 Cancelled { key: K, reason: CancelReason },
101}
102impl<K: PartialEq> Load<K> {
103 pub fn owns(&self, ticket: &Ticket<K>) -> bool {
104 matches!(self, Self::Running(current) if current == ticket)
105 }
106 pub fn covers(&self, key: &K) -> bool {
107 match self {
108 Self::Idle => false,
109 Self::Running(t) => &t.key == key,
110 Self::Ready(k) | Self::Failed { key: k, .. } | Self::Cancelled { key: k, .. } => {
111 k == key
112 }
113 }
114 }
115 pub fn retry_failed(&mut self) {
116 if matches!(self, Self::Failed { .. } | Self::Cancelled { .. }) {
117 *self = Self::Idle;
118 }
119 }
120}
121
122type Resource = Box<dyn FnOnce() -> Result<(), Failure> + Send>;
123#[derive(Default)]
124struct Cancellation {
125 cancelled: AtomicBool,
126 resource: Mutex<Option<Resource>>,
127}
128#[derive(Clone)]
129pub struct CancelToken(Arc<Cancellation>);
130fn invoke(resource: Resource) -> Result<(), Failure> {
131 match std::panic::catch_unwind(std::panic::AssertUnwindSafe(resource)) {
132 Ok(result) => result,
133 Err(_) => Err(Failure::new(
134 FailureKind::Panic,
135 "cancellation resource panicked",
136 )),
137 }
138}
139impl CancelToken {
140 pub fn is_cancelled(&self) -> bool {
141 self.0.cancelled.load(Ordering::Acquire)
142 }
143
144 pub fn register_cancel_resource(
149 &self,
150 resource: impl FnOnce() -> Result<(), Failure> + Send + 'static,
151 ) -> Result<(), Failure> {
152 let resource: Resource = Box::new(resource);
153 {
154 let mut slot = self.0.resource.lock();
155 if !self.is_cancelled() {
156 if slot.is_some() {
157 return Err(Failure::new(
158 FailureKind::Protocol,
159 "cancellation resource already registered",
160 ));
161 }
162 *slot = Some(resource);
163 return Ok(());
164 }
165 }
166 invoke(resource)
167 }
168
169 pub fn clear_cancel_resource(&self) {
172 let resource = self.0.resource.lock().take();
173 drop(resource);
174 }
175
176 fn cancel_resource(&self) -> Result<(), Failure> {
177 let resource = {
178 let mut slot = self.0.resource.lock();
179 self.0.cancelled.store(true, Ordering::Release);
180 slot.take()
181 };
182 match resource {
183 Some(resource) => invoke(resource),
184 None => Ok(()),
185 }
186 }
187}
188type Emitter<T> = Arc<Mutex<Option<Box<dyn FnOnce(Outcome<T>) + Send>>>>;
189fn finish<T>(emitter: &Emitter<T>, outcome: Outcome<T>) {
190 let emit = emitter.lock().take();
191 if let Some(emit) = emit {
192 emit(outcome);
193 }
194}
195pub struct CancelHandle {
196 cancel: Option<Box<dyn FnOnce(CancelReason) + Send>>,
197}
198impl CancelHandle {
199 pub fn cancel(mut self, reason: CancelReason) {
200 if let Some(cancel) = self.cancel.take() {
201 cancel(reason);
202 }
203 }
204}
205impl Drop for CancelHandle {
206 fn drop(&mut self) {
207 if let Some(cancel) = self.cancel.take() {
208 cancel(CancelReason::OwnerClosed);
209 }
210 }
211}
212pub fn spawn<T: Send + 'static>(
213 name: &'static str,
214 emit: impl FnOnce(Outcome<T>) + Send + 'static,
215 work: impl FnOnce(CancelToken) -> Outcome<T> + Send + 'static,
216) -> CancelHandle {
217 let token = CancelToken(Arc::new(Cancellation::default()));
218 let emitter: Emitter<T> = Arc::new(Mutex::new(Some(Box::new(emit))));
219 let cancel_token = token.clone();
220 let cancel_emitter = emitter.clone();
221 let handle = CancelHandle {
222 cancel: Some(Box::new(move |reason| {
223 let emit = cancel_emitter.lock().take();
226 if let Some(emit) = emit {
227 let outcome = match cancel_token.cancel_resource() {
228 Ok(()) => Outcome::Cancelled(reason),
229 Err(failure) => Outcome::Failed {
230 failure,
231 partial: None,
232 },
233 };
234 emit(outcome);
235 }
236 })),
237 };
238 let worker_emitter = emitter.clone();
239 let started = std::thread::Builder::new()
240 .name(name.into())
241 .spawn(move || {
242 if token.is_cancelled() {
243 return;
244 }
245 let outcome = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
246 work(token.clone())
247 })) {
248 Ok(outcome) => outcome,
249 Err(_) => {
250 match token.cancel_resource() {
253 Ok(()) => Outcome::failed(FailureKind::Panic, "worker panicked"),
254 Err(failure) => Outcome::Failed {
255 failure,
256 partial: None,
257 },
258 }
259 }
260 };
261 token.clear_cancel_resource();
262 finish(&worker_emitter, outcome);
263 });
264 if let Err(error) = started {
265 finish(
266 &emitter,
267 Outcome::failed(FailureKind::ThreadStart, error.to_string()),
268 );
269 }
270 handle
271}
272
273#[cfg(test)]
274mod tests {
275 use super::*;
276 use std::sync::mpsc::channel;
277
278 #[test]
279 fn cancellation_reserves_terminal_before_callback_and_worker_success() {
280 let (ready_tx, ready_rx) = channel();
281 let (entered_tx, entered_rx) = channel();
282 let (release_tx, release_rx) = channel();
283 let (work_tx, work_rx) = channel();
284 let (done_tx, done_rx) = channel();
285 let (tx, rx) = channel();
286 let handle = spawn(
287 "race",
288 move |result| {
289 tx.send(result).unwrap();
290 },
291 move |token| {
292 token
293 .register_cancel_resource(move || {
294 entered_tx.send(()).unwrap();
295 release_rx.recv().unwrap();
296 Err(Failure::new(FailureKind::Io, "cleanup failed"))
297 })
298 .unwrap();
299 ready_tx.send(()).unwrap();
300 work_rx.recv().unwrap();
301 done_tx.send(()).unwrap();
302 Outcome::Success(())
303 },
304 );
305 ready_rx.recv().unwrap();
306 let cancel = std::thread::spawn(move || handle.cancel(CancelReason::Dismissed));
307 entered_rx.recv().unwrap();
308 work_tx.send(()).unwrap();
309 done_rx.recv().unwrap();
310 assert!(rx.try_recv().is_err());
311 release_tx.send(()).unwrap();
312 cancel.join().unwrap();
313 assert!(
314 matches!(rx.recv().unwrap(), Outcome::Failed { failure, .. } if failure.kind == FailureKind::Io)
315 );
316 assert!(rx.recv().is_err());
317 }
318
319 #[test]
320 fn late_registration_runs_immediately_and_success_wins_when_already_published() {
321 let token = CancelToken(Arc::new(Cancellation::default()));
322 token.cancel_resource().unwrap();
323 let (tx, rx) = channel();
324 token
325 .register_cancel_resource(move || {
326 tx.send(()).unwrap();
327 Ok(())
328 })
329 .unwrap();
330 rx.recv().unwrap();
331 let (tx, rx) = channel();
332 let handle = spawn(
333 "success",
334 move |result| {
335 tx.send(result).unwrap();
336 },
337 |_| Outcome::Success(7),
338 );
339 assert!(matches!(rx.recv().unwrap(), Outcome::Success(7)));
340 drop(handle);
341 assert!(rx.recv().is_err());
342 }
343}