Skip to main content

ntex_rt/
pool.rs

1//! A thread pool for blocking operations.
2use 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
9/// Submits blocking work and returns a future for its result.
10///
11/// If a system is running, work is submitted to its blocking thread pool.
12/// Otherwise, the closure runs immediately on the current thread.
13///
14/// Dropping the returned future prevents queued work from starting, but cannot
15/// interrupt work that is already running. Call [`BlockingResult::detach`] to
16/// let queued work continue even if its result is no longer needed.
17pub 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/// Error returned when blocking work cannot produce a result.
30///
31/// This can occur if the task is canceled, panics, or no worker thread
32/// can be started.
33#[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/// Future resolving to the result of blocking work.
45#[derive(Debug)]
46pub struct BlockingResult<T> {
47    rx: oneshot::AsyncReceiver<Result<T, Box<dyn Any + Send>>>,
48}
49
50impl<T: 'static> BlockingResult<T> {
51    /// Detaches the task so it can continue without awaiting its result.
52    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
75/// Reserved worker slot, released on drop.
76struct 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                    // release the slot, then pick up work queued meanwhile,
121                    // pairs with the fence in `ThreadPool::execute`
122                    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/// A thread pool for executing blocking operations.
139///
140/// The pool can be configured as either bounded or unbounded, which
141/// determines how tasks are handled when all worker threads are busy.
142///
143/// The number of worker threads scales dynamically with load, but will
144/// never exceed the `thread_limit` parameter. When all worker threads are
145/// busy, tasks are queued until a worker thread becomes available.
146#[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    /// Creates a [`ThreadPool`] with a maximum number of worker threads
160    /// and a timeout for receiving tasks from the task channel.
161    ///
162    /// A `thread_limit` of zero is treated as one.
163    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    /// Submits a closure to the thread pool.
191    ///
192    /// The task will be executed by an available worker thread. If no threads
193    /// are available and the pool has reached its maximum size, the work will
194    /// be queued until a worker thread becomes available. This method never
195    /// blocks.
196    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            // do not execute operation if receiver is dropped
204            if !tx.is_closed() {
205                let result = panic::catch_unwind(panic::AssertUnwindSafe(f));
206                let _ = tx.send(result);
207            }
208        });
209
210        // hand over to an idle worker
211        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        // pairs with the fence in `worker`, either an exiting worker sees
223        // the queued task or a free slot is visible here
224        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                // no worker can run queued tasks, drop them so they
239                // resolve with `BlockingError`
240                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        // concurrent submitters
297        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            // submit around the moment the idle worker times out
330            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}