1use std::{
4 cell::RefCell,
5 sync::{
6 atomic::{AtomicBool, AtomicU8, Ordering},
7 Arc,
8 },
9};
10
11const RUNNING: u8 = 0;
13const FINALIZATION_REQUESTED: u8 = 1;
15const CANCELLED: u8 = 2;
17
18#[derive(Default, Clone, Debug)]
25pub struct CancelOnDrop(Arc<AtomicU8>);
26
27impl CancelOnDrop {
30 pub fn is_interrupted(&self) -> bool {
32 self.0.load(Ordering::Relaxed) != RUNNING
33 }
34
35 pub fn is_cancelled(&self) -> bool {
37 self.0.load(Ordering::Relaxed) == CANCELLED
38 }
39
40 pub fn request_finalization(&self) {
42 let _ = self.0.compare_exchange(
43 RUNNING,
44 FINALIZATION_REQUESTED,
45 Ordering::Relaxed,
46 Ordering::Relaxed,
47 );
48 }
49
50 pub fn is_finalization_requested(&self) -> bool {
52 self.0.load(Ordering::Relaxed) == FINALIZATION_REQUESTED
53 }
54
55 pub fn scope<R>(&self, f: impl FnOnce() -> R) -> R {
61 struct Restore(Option<Arc<AtomicU8>>);
62
63 impl Drop for Restore {
64 fn drop(&mut self) {
65 CURRENT.set(self.0.take());
66 }
67 }
68
69 let _restore = Restore(CURRENT.replace(Some(self.0.clone())));
70 f()
71 }
72}
73
74impl Drop for CancelOnDrop {
75 fn drop(&mut self) {
76 self.0.store(CANCELLED, Ordering::Relaxed);
77 }
78}
79
80#[derive(Default, Clone, Debug)]
87pub struct ManualCancel(Arc<AtomicBool>);
88
89impl ManualCancel {
92 pub fn is_cancelled(&self) -> bool {
94 self.0.load(Ordering::Relaxed)
95 }
96
97 pub fn cancel(self) {
99 self.0.store(true, Ordering::Relaxed);
100 }
101}
102
103pub fn is_cancelled() -> bool {
112 CURRENT.with_borrow(|state| {
113 state.as_ref().is_some_and(|state| state.load(Ordering::Relaxed) == CANCELLED)
114 })
115}
116
117thread_local! {
118 static CURRENT: RefCell<Option<Arc<AtomicU8>>> = const { RefCell::new(None) };
120}
121
122#[cfg(test)]
123mod tests {
124 use super::*;
125
126 #[test]
127 fn test_default_cancelled() {
128 let c = CancelOnDrop::default();
129 assert!(!c.is_interrupted());
130 assert!(!c.is_cancelled());
131 }
132
133 #[test]
134 fn test_default_cancel_task() {
135 let c = ManualCancel::default();
136 assert!(!c.is_cancelled());
137 }
138
139 #[test]
140 fn test_set_cancel_task() {
141 let c = ManualCancel::default();
142 assert!(!c.is_cancelled());
143 let c2 = c.clone();
144 let c3 = c.clone();
145 c.cancel();
146 assert!(c3.is_cancelled());
147 assert!(c2.is_cancelled());
148 }
149
150 #[test]
151 fn test_cancel_task_multiple_threads() {
152 let c = ManualCancel::default();
153 let cloned_cancel = c.clone();
154
155 let mut handles = vec![];
159 for _ in 0..10 {
160 let c = c.clone();
161 let handle = std::thread::spawn(move || {
162 for _ in 0..1000 {
163 if c.is_cancelled() {
164 return;
165 }
166 }
167 });
168 handles.push(handle);
169 }
170
171 for handle in handles {
173 handle.join().unwrap();
174 }
175
176 assert!(!c.is_cancelled());
178
179 c.cancel();
181 assert!(cloned_cancel.is_cancelled());
182 }
183
184 #[test]
185 fn test_cancelondrop_clone_behavior() {
186 let cancel = CancelOnDrop::default();
187 assert!(!cancel.is_cancelled());
188
189 let cloned_cancel = cancel.clone();
191 assert!(!cloned_cancel.is_cancelled());
192
193 drop(cancel);
195
196 assert!(cloned_cancel.is_interrupted());
198 assert!(cloned_cancel.is_cancelled());
199 }
200
201 #[test]
202 fn test_cancelondrop_multiple_clones() {
203 let cancel = CancelOnDrop::default();
204 let clone1 = cancel.clone();
205 let clone2 = cancel.clone();
206 let clone3 = cancel.clone();
207
208 assert!(!cancel.is_cancelled());
209 assert!(!clone1.is_cancelled());
210 assert!(!clone2.is_cancelled());
211 assert!(!clone3.is_cancelled());
212
213 drop(clone1);
215
216 assert!(cancel.is_interrupted());
217 assert!(cancel.is_cancelled());
218 assert!(clone2.is_cancelled());
219 assert!(clone3.is_cancelled());
220 }
221
222 #[test]
223 fn test_cancel_on_drop_finalization_request() {
224 let cancel = CancelOnDrop::default();
225 let clone = cancel.clone();
226
227 cancel.request_finalization();
228
229 assert!(clone.is_interrupted());
230 assert!(clone.is_finalization_requested());
231 assert!(!clone.is_cancelled());
232
233 drop(cancel);
234
235 assert!(clone.is_cancelled());
236 assert!(!clone.is_finalization_requested());
237
238 clone.request_finalization();
239
240 assert!(clone.is_cancelled());
241 assert!(!clone.is_finalization_requested());
242 }
243
244 #[test]
245 fn test_scope_is_cancelled() {
246 assert!(!is_cancelled());
247
248 let outer = CancelOnDrop::default();
249 let inner = CancelOnDrop::default();
250 outer.scope(|| {
251 assert!(!is_cancelled());
252 inner.scope(|| {
253 drop(inner.clone());
254 assert!(is_cancelled());
255 });
256 assert!(!is_cancelled());
258 outer.request_finalization();
259 assert!(!is_cancelled());
260 });
261
262 assert!(!is_cancelled());
264 assert!(!outer.is_cancelled());
265 }
266
267 #[test]
268 fn test_scope_restores_on_panic() {
269 let cancel = CancelOnDrop::default();
270 drop(cancel.clone());
271
272 let res = std::panic::catch_unwind(|| cancel.scope(|| panic!("scope panicked")));
273 assert!(res.is_err());
274 assert!(!is_cancelled());
275 }
276}