1use std::sync::Arc;
3use std::sync::atomic::{AtomicUsize, Ordering, fence};
4use std::task::{Context, Poll};
5use std::{any::Any, fmt, future::Future, panic, pin::Pin, thread, time::Duration};
6
7use crossbeam_channel::{Receiver, Select, Sender, TrySendError, bounded, unbounded};
8
9pub fn spawn_blocking<F, R>(f: F) -> BlockingResult<R>
18where
19 F: FnOnce() -> R + Send + 'static,
20 R: Send + 'static,
21{
22 if let Some(sys) = crate::System::try_current() {
23 sys.spawn_blocking(f)
24 } else {
25 ThreadPool::execute_inplace(f)
26 }
27}
28
29#[derive(Copy, Clone, Debug, PartialEq, Eq)]
34pub struct BlockingError;
35
36impl std::error::Error for BlockingError {}
37
38impl fmt::Display for BlockingError {
39 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
40 "Blocking task failed or was canceled".fmt(f)
41 }
42}
43
44#[derive(Debug)]
46pub struct BlockingResult<T> {
47 rx: oneshot::AsyncReceiver<Result<T, Box<dyn Any + Send>>>,
48}
49
50impl<T: 'static> BlockingResult<T> {
51 pub fn detach(self) {
53 crate::spawn(async move {
54 let _ = self.await;
55 })
56 .detach();
57 }
58}
59
60type BoxedDispatchable = Box<dyn Dispatchable + Send>;
61
62pub(crate) trait Dispatchable: Send + 'static {
63 fn run(self: Box<Self>);
64}
65
66impl<F> Dispatchable for F
67where
68 F: FnOnce() + Send + 'static,
69{
70 fn run(self: Box<Self>) {
71 (*self)();
72 }
73}
74
75struct CounterGuard(Arc<AtomicUsize>);
77
78impl CounterGuard {
79 fn reserve(counter: &Arc<AtomicUsize>, limit: usize) -> Option<(Self, usize)> {
80 counter
81 .try_update(Ordering::AcqRel, Ordering::Acquire, |cnt| {
82 (cnt < limit).then_some(cnt + 1)
83 })
84 .ok()
85 .map(|cnt| (CounterGuard(counter.clone()), cnt))
86 }
87}
88
89impl Drop for CounterGuard {
90 fn drop(&mut self) {
91 self.0.fetch_sub(1, Ordering::AcqRel);
92 }
93}
94
95fn worker(
96 receiver_high_prio: Receiver<BoxedDispatchable>,
97 receiver_low_prio: Receiver<BoxedDispatchable>,
98 guard: CounterGuard,
99 thread_limit: usize,
100 timeout: Duration,
101) -> impl FnOnce() {
102 move || {
103 let mut guard = guard;
104 let mut sel = Select::new_biased();
105 sel.recv(&receiver_high_prio);
106 sel.recv(&receiver_low_prio);
107 loop {
108 match sel.select_timeout(timeout) {
109 Ok(op) if op.index() == 0 => {
110 if let Ok(f) = op.recv(&receiver_high_prio) {
111 f.run();
112 }
113 }
114 Ok(op) => {
115 if let Ok(f) = op.recv(&receiver_low_prio) {
116 f.run();
117 }
118 }
119 Err(_) => {
120 let counter = guard.0.clone();
123 drop(guard);
124 fence(Ordering::SeqCst);
125 if receiver_high_prio.is_empty() {
126 return;
127 }
128 match CounterGuard::reserve(&counter, thread_limit) {
129 Some((g, _)) => guard = g,
130 None => return,
131 }
132 }
133 }
134 }
135 }
136}
137
138#[derive(Debug, Clone)]
147pub struct ThreadPool {
148 name: String,
149 sender_low_prio: Sender<BoxedDispatchable>,
150 receiver_low_prio: Receiver<BoxedDispatchable>,
151 sender_high_prio: Sender<BoxedDispatchable>,
152 receiver_high_prio: Receiver<BoxedDispatchable>,
153 counter: Arc<AtomicUsize>,
154 thread_limit: usize,
155 recv_timeout: Duration,
156}
157
158impl ThreadPool {
159 pub fn new(name: &str, thread_limit: usize, recv_timeout: Duration) -> Self {
164 let (sender_low_prio, receiver_low_prio) = bounded(0);
165 let (sender_high_prio, receiver_high_prio) = unbounded();
166 Self {
167 sender_low_prio,
168 receiver_low_prio,
169 sender_high_prio,
170 receiver_high_prio,
171 thread_limit: thread_limit.max(1),
172 recv_timeout,
173 name: format!("{name}:pool-wrk"),
174 counter: Arc::new(AtomicUsize::new(0)),
175 }
176 }
177
178 pub(crate) fn execute_inplace<F, R>(f: F) -> BlockingResult<R>
179 where
180 F: FnOnce() -> R + Send + 'static,
181 R: Send + 'static,
182 {
183 let (tx, rx) = oneshot::async_channel();
184 let result = panic::catch_unwind(panic::AssertUnwindSafe(f));
185 let _ = tx.send(result);
186 BlockingResult { rx }
187 }
188
189 #[allow(clippy::missing_panics_doc)]
190 pub fn execute<F, R>(&self, f: F) -> BlockingResult<R>
197 where
198 F: FnOnce() -> R + Send + 'static,
199 R: Send + 'static,
200 {
201 let (tx, rx) = oneshot::async_channel();
202 let f = Box::new(move || {
203 if !tx.is_closed() {
205 let result = panic::catch_unwind(panic::AssertUnwindSafe(f));
206 let _ = tx.send(result);
207 }
208 });
209
210 let f = match self.sender_low_prio.try_send(f) {
212 Ok(()) => return BlockingResult { rx },
213 Err(TrySendError::Full(f)) => f,
214 Err(TrySendError::Disconnected(_)) => {
215 unreachable!("receiver should not all disconnected")
216 }
217 };
218
219 self.sender_high_prio
220 .send(f)
221 .expect("the channel should not be closed");
222 fence(Ordering::SeqCst);
225
226 if let Some((guard, idx)) = CounterGuard::reserve(&self.counter, self.thread_limit) {
227 let result = thread::Builder::new()
228 .name(format!("{}:{}", self.name, idx))
229 .spawn(worker(
230 self.receiver_high_prio.clone(),
231 self.receiver_low_prio.clone(),
232 guard,
233 self.thread_limit,
234 self.recv_timeout,
235 ));
236 if let Err(e) = result {
237 log::error!("Cannot start blocking pool thread: {e}");
238 while self.counter.load(Ordering::Acquire) == 0 {
241 if self.receiver_high_prio.try_recv().is_err() {
242 break;
243 }
244 }
245 }
246 }
247 BlockingResult { rx }
248 }
249}
250
251impl<R> Future for BlockingResult<R> {
252 type Output = Result<R, BlockingError>;
253
254 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
255 let this = self.get_mut();
256
257 match Pin::new(&mut this.rx).poll(cx) {
258 Poll::Pending => Poll::Pending,
259 Poll::Ready(result) => Poll::Ready(
260 result
261 .map_err(|_| BlockingError)
262 .and_then(|res| res.map_err(|_| BlockingError)),
263 ),
264 }
265 }
266}
267
268#[cfg(test)]
269mod tests {
270 use std::time::Instant;
271
272 use super::*;
273
274 fn wait<R>(fut: BlockingResult<R>, timeout: Duration) -> Option<Result<R, BlockingError>> {
275 let mut fut = std::pin::pin!(fut);
276 let mut cx = Context::from_waker(std::task::Waker::noop());
277 let start = Instant::now();
278 loop {
279 if let Poll::Ready(res) = fut.as_mut().poll(&mut cx) {
280 return Some(res);
281 }
282 if start.elapsed() > timeout {
283 return None;
284 }
285 thread::sleep(Duration::from_millis(1));
286 }
287 }
288
289 #[test]
290 fn thread_limit_respected() {
291 let pool = ThreadPool::new("test", 2, Duration::from_secs(1));
292 let running = Arc::new(AtomicUsize::new(0));
293 let max = Arc::new(AtomicUsize::new(0));
294 let barrier = Arc::new(std::sync::Barrier::new(8));
295
296 let submitters: Vec<_> = (0..8)
298 .map(|_| {
299 let (pool, running, max, barrier) =
300 (pool.clone(), running.clone(), max.clone(), barrier.clone());
301 thread::spawn(move || {
302 barrier.wait();
303 (0..8)
304 .map(|_| {
305 let (running, max) = (running.clone(), max.clone());
306 pool.execute(move || {
307 let cnt = running.fetch_add(1, Ordering::SeqCst) + 1;
308 max.fetch_max(cnt, Ordering::SeqCst);
309 thread::sleep(Duration::from_millis(5));
310 running.fetch_sub(1, Ordering::SeqCst);
311 })
312 })
313 .collect::<Vec<_>>()
314 })
315 })
316 .collect();
317 for s in submitters {
318 for res in s.join().unwrap() {
319 assert_eq!(wait(res, Duration::from_secs(10)), Some(Ok(())));
320 }
321 }
322 assert!(max.load(Ordering::SeqCst) <= 2, "{max:?} tasks ran at once");
323 }
324
325 #[test]
326 fn idle_workers_do_not_strand_tasks() {
327 let pool = ThreadPool::new("test", 1, Duration::from_millis(2));
328 for i in 0..300u64 {
329 thread::sleep(Duration::from_micros(1500 + (i % 10) * 100));
331 let res = wait(pool.execute(move || i), Duration::from_secs(5));
332 assert_eq!(res, Some(Ok(i)), "task {i} was not executed");
333 }
334 }
335
336 #[test]
337 fn spawn_blocking_without_system() {
338 thread::spawn(|| {
339 let tid = thread::current().id();
340 let res = spawn_blocking(move || thread::current().id() == tid);
341 assert_eq!(wait(res, Duration::from_secs(1)), Some(Ok(true)));
342
343 let res = spawn_blocking(|| panic!("blocking"));
344 assert_eq!(
345 wait(res, Duration::from_secs(1)),
346 Some(Err::<(), _>(BlockingError))
347 );
348 })
349 .join()
350 .unwrap();
351 assert_eq!(
352 BlockingError.to_string(),
353 "Blocking task failed or was canceled"
354 );
355 }
356
357 #[test]
358 fn detached_blocking_task_runs() {
359 crate::System::new("test", crate::testing::TestRunner).block_on(async {
360 let (tx, rx) = oneshot::async_channel();
361 spawn_blocking(move || tx.send(1).unwrap()).detach();
362 assert_eq!(rx.await, Ok(1));
363 });
364 }
365
366 #[test]
367 fn zero_thread_limit() {
368 let pool = ThreadPool::new("test", 0, Duration::from_secs(1));
369 assert_eq!(
370 wait(pool.execute(|| 1), Duration::from_secs(5)),
371 Some(Ok(1))
372 );
373 }
374}