1use alloc::{boxed::Box, format, string::String, vec::Vec};
2use core::{future::Future, marker::PhantomData, mem, panic::AssertUnwindSafe};
3use std::{
4 thread::{self, JoinHandle},
5thread_local,
6};
78use crate::executor::FallibleTask;
9use bevy_platform::sync::Arc;
10use concurrent_queue::ConcurrentQueue;
11use futures_lite::FutureExt;
1213use crate::{
14block_on,
15 thread_executor::{ThreadExecutor, ThreadExecutorTicker},
16Task,
17};
1819struct CallOnDrop(Option<Arc<dyn Fn() + Send + Sync + 'static>>);
2021impl Dropfor CallOnDrop {
22fn drop(&mut self) {
23if let Some(call) = self.0.as_ref() {
24call();
25 }
26 }
27}
2829/// Used to create a [`TaskPool`]
30#[derive(#[automatically_derived]
impl ::core::default::Default for TaskPoolBuilder {
#[inline]
fn default() -> Self {
Self {
num_threads: ::core::default::Default::default(),
stack_size: ::core::default::Default::default(),
thread_name: ::core::default::Default::default(),
on_thread_spawn: ::core::default::Default::default(),
on_thread_destroy: ::core::default::Default::default(),
}
}
}Default)]
31#[must_use]
32pub struct TaskPoolBuilder {
33/// If set, we'll set up the thread pool to use at most `num_threads` threads.
34 /// Otherwise use the logical core count of the system
35num_threads: Option<usize>,
36/// If set, we'll use the given stack size rather than the system default
37stack_size: Option<usize>,
38/// Allows customizing the name of the threads - helpful for debugging. If set, threads will
39 /// be named `<thread_name> (<thread_index>)`, i.e. `"MyThreadPool (2)"`.
40thread_name: Option<String>,
4142 on_thread_spawn: Option<Arc<dyn Fn() + Send + Sync + 'static>>,
43 on_thread_destroy: Option<Arc<dyn Fn() + Send + Sync + 'static>>,
44}
4546impl TaskPoolBuilder {
47/// Creates a new [`TaskPoolBuilder`] instance
48pub fn new() -> Self {
49Self::default()
50 }
5152/// Override the number of threads created for the pool. If unset, we default to the number
53 /// of logical cores of the system
54pub fn num_threads(mut self, num_threads: usize) -> Self {
55self.num_threads = Some(num_threads);
56self57 }
5859/// Override the stack size of the threads created for the pool
60pub fn stack_size(mut self, stack_size: usize) -> Self {
61self.stack_size = Some(stack_size);
62self63 }
6465/// Override the name of the threads created for the pool. If set, threads will
66 /// be named `<thread_name> (<thread_index>)`, i.e. `MyThreadPool (2)`
67pub fn thread_name(mut self, thread_name: String) -> Self {
68self.thread_name = Some(thread_name);
69self70 }
7172/// Sets a callback that is invoked once for every created thread as it starts.
73 ///
74 /// This is called on the thread itself and has access to all thread-local storage.
75 /// This will block running async tasks on the thread until the callback completes.
76pub fn on_thread_spawn(mut self, f: impl Fn() + Send + Sync + 'static) -> Self {
77let arc = Arc::new(f);
7879#[cfg(not(target_has_atomic = "ptr"))]
80 #[expect(
81 unsafe_code,
82 reason = "unsized coercion is an unstable feature for non-std types"
83)]
84// SAFETY:
85 // - Coercion from `impl Fn` to `dyn Fn` is valid
86 // - `Arc::from_raw` receives a valid pointer from a previous call to `Arc::into_raw`
87let arc = unsafe {
88 Arc::from_raw(Arc::into_raw(arc) as *const (dyn Fn() + Send + Sync + 'static))
89 };
9091self.on_thread_spawn = Some(arc);
92self93 }
9495/// Sets a callback that is invoked once for every created thread as it terminates.
96 ///
97 /// This is called on the thread itself and has access to all thread-local storage.
98 /// This will block thread termination until the callback completes.
99pub fn on_thread_destroy(mut self, f: impl Fn() + Send + Sync + 'static) -> Self {
100let arc = Arc::new(f);
101102#[cfg(not(target_has_atomic = "ptr"))]
103 #[expect(
104 unsafe_code,
105 reason = "unsized coercion is an unstable feature for non-std types"
106)]
107// SAFETY:
108 // - Coercion from `impl Fn` to `dyn Fn` is valid
109 // - `Arc::from_raw` receives a valid pointer from a previous call to `Arc::into_raw`
110let arc = unsafe {
111 Arc::from_raw(Arc::into_raw(arc) as *const (dyn Fn() + Send + Sync + 'static))
112 };
113114self.on_thread_destroy = Some(arc);
115self116 }
117118/// Creates a new [`TaskPool`] based on the current options.
119pub fn build(self) -> TaskPool {
120TaskPool::new_internal(self)
121 }
122}
123124/// A thread pool for executing tasks.
125///
126/// While futures usually need to be polled to be executed, Bevy tasks are being
127/// automatically driven by the pool on threads owned by the pool. The [`Task`]
128/// future only needs to be polled in order to receive the result. (For that
129/// purpose, it is often stored in a component or resource, see the
130/// `async_compute` example.)
131///
132/// If the result is not required, one may also use [`Task::detach`] and the pool
133/// will still execute a task, even if it is dropped.
134#[derive(#[automatically_derived]
impl ::core::fmt::Debug for TaskPool {
#[inline]
fn fmt(&self, f: &mut ::core::fmt::Formatter) -> ::core::fmt::Result {
::core::fmt::Formatter::debug_struct_field3_finish(f, "TaskPool",
"executor", &self.executor, "threads", &self.threads,
"shutdown_tx", &&self.shutdown_tx)
}
}Debug)]
135pub struct TaskPool {
136/// The executor for the pool.
137executor: Arc<crate::executor::Executor<'static>>,
138139// The inner state of the pool.
140threads: Vec<JoinHandle<()>>,
141 shutdown_tx: async_channel::Sender<()>,
142}
143144impl TaskPool {
145const LOCAL_EXECUTOR:
::std::thread::LocalKey<crate::executor::LocalExecutor<'static>> =
{
const __RUST_STD_INTERNAL_INIT:
crate::executor::LocalExecutor<'static> =
{ crate::executor::LocalExecutor::new() };
unsafe {
::std::thread::LocalKey::new(const {
if ::std::mem::needs_drop::<crate::executor::LocalExecutor<'static>>()
{
|_|
{
#[thread_local]
static __RUST_STD_INTERNAL_VAL:
::std::thread::local_impl::EagerStorage<crate::executor::LocalExecutor<'static>>
=
::std::thread::local_impl::EagerStorage::new(__RUST_STD_INTERNAL_INIT);
__RUST_STD_INTERNAL_VAL.get()
}
} else {
|_|
{
#[thread_local]
static __RUST_STD_INTERNAL_VAL:
crate::executor::LocalExecutor<'static> =
__RUST_STD_INTERNAL_INIT;
&__RUST_STD_INTERNAL_VAL
}
}
})
}
};
const THREAD_EXECUTOR: ::std::thread::LocalKey<Arc<ThreadExecutor<'static>>> =
{
#[allow(mismatched_lifetime_syntaxes)]
#[inline]
fn __rust_std_internal_init_fn(_lifetime_elision:
::std::marker::PhantomData<&'static ()>)
-> Arc<ThreadExecutor<'static>> {
Arc::new(ThreadExecutor::new())
}
unsafe {
::std::thread::LocalKey::new(const {
if ::std::mem::needs_drop::<Arc<ThreadExecutor<'static>>>()
{
|__rust_std_internal_init|
{
#[thread_local]
static __RUST_STD_INTERNAL_VAL:
::std::thread::local_impl::LazyStorage<Arc<ThreadExecutor<'static>>,
()> =
::std::thread::local_impl::LazyStorage::new();
__RUST_STD_INTERNAL_VAL.get_or_init(__rust_std_internal_init,
|| __rust_std_internal_init_fn(::std::marker::PhantomData))
}
} else {
|__rust_std_internal_init|
{
#[thread_local]
static __RUST_STD_INTERNAL_VAL:
::std::thread::local_impl::LazyStorage<Arc<ThreadExecutor<'static>>,
!> =
::std::thread::local_impl::LazyStorage::new();
__RUST_STD_INTERNAL_VAL.get_or_init(__rust_std_internal_init,
|| __rust_std_internal_init_fn(::std::marker::PhantomData))
}
}
})
}
};thread_local! {
146static LOCAL_EXECUTOR: crate::executor::LocalExecutor<'static> = const { crate::executor::LocalExecutor::new() };
147static THREAD_EXECUTOR: Arc<ThreadExecutor<'static>> = Arc::new(ThreadExecutor::new());
148 }149150/// Each thread should only create one `ThreadExecutor`, otherwise, there are good chances they will deadlock
151pub fn get_thread_executor() -> Arc<ThreadExecutor<'static>> {
152Self::THREAD_EXECUTOR.with(Clone::clone)
153 }
154155/// Create a `TaskPool` with the default configuration.
156pub fn new() -> Self {
157TaskPoolBuilder::new().build()
158 }
159160fn new_internal(builder: TaskPoolBuilder) -> Self {
161let (shutdown_tx, shutdown_rx) = async_channel::unbounded::<()>();
162163let executor = Arc::new(crate::executor::Executor::new());
164165let num_requested_threads = builder166 .num_threads
167 .unwrap_or_else(crate::available_parallelism);
168169// In tests, we want there to be at least two threads so that we're
170 // actually testing multithreaded behavior.
171#[cfg(all(test, feature = "multi_threaded"))]
172let num_threads = num_requested_threads.max(2);
173#[cfg(not(all(test, feature = "multi_threaded")))]
174let num_threads = num_requested_threads;
175176let threads = (0..num_threads)
177 .map(|i| {
178let ex = Arc::clone(&executor);
179let shutdown_rx = shutdown_rx.clone();
180181let thread_name = if let Some(thread_name) = builder.thread_name.as_deref() {
182::alloc::__export::must_use({
::alloc::fmt::format(format_args!("{0} ({1})", thread_name, i))
})format!("{thread_name} ({i})")183 } else {
184::alloc::__export::must_use({
::alloc::fmt::format(format_args!("TaskPool ({0})", i))
})format!("TaskPool ({i})")185 };
186let mut thread_builder = thread::Builder::new().name(thread_name);
187188if let Some(stack_size) = builder.stack_size {
189thread_builder = thread_builder.stack_size(stack_size);
190 }
191192let on_thread_spawn = builder.on_thread_spawn.clone();
193let on_thread_destroy = builder.on_thread_destroy.clone();
194195thread_builder196 .spawn(move || {
197TaskPool::LOCAL_EXECUTOR.with(|local_executor| {
198if let Some(on_thread_spawn) = on_thread_spawn {
199on_thread_spawn();
200drop(on_thread_spawn);
201 }
202let _destructor = CallOnDrop(on_thread_destroy);
203loop {
204let res = std::panic::catch_unwind(|| {
205let tick_forever = async move {
206loop {
207 local_executor.tick().await;
208 }
209 };
210block_on(ex.run(tick_forever.or(shutdown_rx.recv())))
211 });
212if let Ok(value) = res {
213// Use unwrap_err because we expect a Closed error
214value.unwrap_err();
215break;
216 }
217 }
218 });
219 })
220 .expect("Failed to spawn thread.")
221 })
222 .collect();
223224Self {
225executor,
226threads,
227shutdown_tx,
228 }
229 }
230231/// Return the number of threads owned by the task pool
232pub fn thread_num(&self) -> usize {
233self.threads.len()
234 }
235236/// Allows spawning non-`'static` futures on the thread pool. The function takes a callback,
237 /// passing a scope object into it. The scope object provided to the callback can be used
238 /// to spawn tasks. This function will await the completion of all tasks before returning.
239 ///
240 /// This is similar to [`thread::scope`] and `rayon::scope`.
241 ///
242 /// # Example
243 ///
244 /// ```
245 /// use bevy_tasks::TaskPool;
246 ///
247 /// let pool = TaskPool::new();
248 /// let mut x = 0;
249 /// let results = pool.scope(|s| {
250 /// s.spawn(async {
251 /// // you can borrow the spawner inside a task and spawn tasks from within the task
252 /// s.spawn(async {
253 /// // borrow x and mutate it.
254 /// x = 2;
255 /// // return a value from the task
256 /// 1
257 /// });
258 /// // return some other value from the first task
259 /// 0
260 /// });
261 /// });
262 ///
263 /// // The ordering of results is non-deterministic if you spawn from within tasks as above.
264 /// // If you're doing this, you'll have to write your code to not depend on the ordering.
265 /// assert!(results.contains(&0));
266 /// assert!(results.contains(&1));
267 ///
268 /// // The ordering is deterministic if you only spawn directly from the closure function.
269 /// let results = pool.scope(|s| {
270 /// s.spawn(async { 0 });
271 /// s.spawn(async { 1 });
272 /// });
273 /// assert_eq!(&results[..], &[0, 1]);
274 ///
275 /// // You can access x after scope runs, since it was only temporarily borrowed in the scope.
276 /// assert_eq!(x, 2);
277 /// ```
278 ///
279 /// # Lifetimes
280 ///
281 /// The [`Scope`] object takes two lifetimes: `'scope` and `'env`.
282 ///
283 /// The `'scope` lifetime represents the lifetime of the scope. That is the time during
284 /// which the provided closure and tasks that are spawned into the scope are run.
285 ///
286 /// The `'env` lifetime represents the lifetime of whatever is borrowed by the scope.
287 /// Thus this lifetime must outlive `'scope`.
288 ///
289 /// ```compile_fail
290 /// use bevy_tasks::TaskPool;
291 /// fn scope_escapes_closure() {
292 /// let pool = TaskPool::new();
293 /// let foo = Box::new(42);
294 /// pool.scope(|scope| {
295 /// std::thread::spawn(move || {
296 /// // UB. This could spawn on the scope after `.scope` returns and the internal Scope is dropped.
297 /// scope.spawn(async move {
298 /// assert_eq!(*foo, 42);
299 /// });
300 /// });
301 /// });
302 /// }
303 /// ```
304 ///
305 /// ```compile_fail
306 /// use bevy_tasks::TaskPool;
307 /// fn cannot_borrow_from_closure() {
308 /// let pool = TaskPool::new();
309 /// pool.scope(|scope| {
310 /// let x = 1;
311 /// let y = &x;
312 /// scope.spawn(async move {
313 /// assert_eq!(*y, 1);
314 /// });
315 /// });
316 /// }
317pub fn scope<'env, F, T>(&self, f: F) -> Vec<T>
318where
319F: for<'scope> FnOnce(&'scope Scope<'scope, 'env, T>),
320 T: Send + 'static,
321 {
322Self::THREAD_EXECUTOR.with(|scope_executor| {
323self.scope_with_executor_inner(true, scope_executor, scope_executor, f)
324 })
325 }
326327/// This allows passing an external executor to spawn tasks on. When you pass an external executor
328 /// [`Scope::spawn_on_scope`] spawns is then run on the thread that [`ThreadExecutor`] is being ticked on.
329 /// If [`None`] is passed the scope will use a [`ThreadExecutor`] that is ticked on the current thread.
330 ///
331 /// When `tick_task_pool_executor` is set to `true`, the multithreaded task stealing executor is ticked on the scope
332 /// thread. Disabling this can be useful when finishing the scope is latency sensitive. Pulling tasks from
333 /// global executor can run tasks unrelated to the scope and delay when the scope returns.
334 ///
335 /// See [`Self::scope`] for more details in general about how scopes work.
336pub fn scope_with_executor<'env, F, T>(
337&self,
338 tick_task_pool_executor: bool,
339 external_executor: Option<&ThreadExecutor>,
340 f: F,
341 ) -> Vec<T>
342where
343F: for<'scope> FnOnce(&'scope Scope<'scope, 'env, T>),
344 T: Send + 'static,
345 {
346Self::THREAD_EXECUTOR.with(|scope_executor| {
347// If an `external_executor` is passed, use that. Otherwise, get the executor stored
348 // in the `THREAD_EXECUTOR` thread local.
349if let Some(external_executor) = external_executor {
350self.scope_with_executor_inner(
351tick_task_pool_executor,
352external_executor,
353scope_executor,
354f,
355 )
356 } else {
357self.scope_with_executor_inner(
358tick_task_pool_executor,
359scope_executor,
360scope_executor,
361f,
362 )
363 }
364 })
365 }
366367#[expect(unsafe_code, reason = "Required to transmute lifetimes.")]
368fn scope_with_executor_inner<'env, F, T>(
369&self,
370 tick_task_pool_executor: bool,
371 external_executor: &ThreadExecutor,
372 scope_executor: &ThreadExecutor,
373 f: F,
374 ) -> Vec<T>
375where
376F: for<'scope> FnOnce(&'scope Scope<'scope, 'env, T>),
377 T: Send + 'static,
378 {
379// SAFETY: This safety comment applies to all references transmuted to 'env.
380 // Any futures spawned with these references need to return before this function completes.
381 // This is guaranteed because we drive all the futures spawned onto the Scope
382 // to completion in this function. However, rust has no way of knowing this so we
383 // transmute the lifetimes to 'env here to appease the compiler as it is unable to validate safety.
384 // Any usages of the references passed into `Scope` must be accessed through
385 // the transmuted reference for the rest of this function.
386let executor: &crate::executor::Executor = &self.executor;
387// SAFETY: As above, all futures must complete in this function so we can change the lifetime
388let executor: &'env crate::executor::Executor = unsafe { mem::transmute(executor) };
389// SAFETY: As above, all futures must complete in this function so we can change the lifetime
390let external_executor: &'env ThreadExecutor<'env> =
391unsafe { mem::transmute(external_executor) };
392// SAFETY: As above, all futures must complete in this function so we can change the lifetime
393let scope_executor: &'env ThreadExecutor<'env> = unsafe { mem::transmute(scope_executor) };
394let spawned: ConcurrentQueue<FallibleTask<Result<T, Box<dyn core::any::Any + Send>>>> =
395ConcurrentQueue::unbounded();
396// shadow the variable so that the owned value cannot be used for the rest of the function
397 // SAFETY: As above, all futures must complete in this function so we can change the lifetime
398let spawned: &'env ConcurrentQueue<
399FallibleTask<Result<T, Box<dyn core::any::Any + Send>>>,
400 > = unsafe { mem::transmute(&spawned) };
401402let scope = Scope {
403executor,
404external_executor,
405scope_executor,
406spawned,
407 scope: PhantomData,
408 env: PhantomData,
409 };
410411// shadow the variable so that the owned value cannot be used for the rest of the function
412 // SAFETY: As above, all futures must complete in this function so we can change the lifetime
413let scope: &'env Scope<'_, 'env, T> = unsafe { mem::transmute(&scope) };
414415f(scope);
416417if spawned.is_empty() {
418Vec::new()
419 } else {
420block_on(async move {
421let get_results = async {
422let mut results = Vec::with_capacity(spawned.len());
423while let Ok(task) = spawned.pop() {
424if let Some(res) = task.await {
425match res {
426Ok(res) => results.push(res),
427Err(payload) => std::panic::resume_unwind(payload),
428 }
429 } else {
430{ ::core::panicking::panic_fmt(format_args!("Failed to catch panic!")); };panic!("Failed to catch panic!");
431 }
432 }
433results434 };
435436let tick_task_pool_executor = tick_task_pool_executor || self.threads.is_empty();
437438// we get this from a thread local so we should always be on the scope executors thread.
439 // note: it is possible `scope_executor` and `external_executor` is the same executor,
440 // in that case, we should only tick one of them, otherwise, it may cause deadlock.
441let scope_ticker = scope_executor.ticker().unwrap();
442let external_ticker = if !external_executor.is_same(scope_executor) {
443external_executor.ticker()
444 } else {
445None446 };
447448match (external_ticker, tick_task_pool_executor) {
449 (Some(external_ticker), true) => {
450Self::execute_global_external_scope(
451 executor,
452 external_ticker,
453 scope_ticker,
454 get_results,
455 )
456 .await
457}
458 (Some(external_ticker), false) => {
459Self::execute_external_scope(external_ticker, scope_ticker, get_results)
460 .await
461}
462// either external_executor is none or it is same as scope_executor
463(None, true) => {
464Self::execute_global_scope(executor, scope_ticker, get_results).await
465}
466 (None, false) => Self::execute_scope(scope_ticker, get_results).await,
467 }
468 })
469 }
470 }
471472#[inline]
473async fn execute_global_external_scope<'scope, 'ticker, T>(
474 executor: &'scope crate::executor::Executor<'scope>,
475 external_ticker: ThreadExecutorTicker<'scope, 'ticker>,
476 scope_ticker: ThreadExecutorTicker<'scope, 'ticker>,
477 get_results: impl Future<Output = Vec<T>>,
478 ) -> Vec<T> {
479// we restart the executors if a task errors. if a scoped
480 // task errors it will panic the scope on the call to get_results
481let execute_forever = async move {
482loop {
483let tick_forever = async {
484loop {
485 external_ticker.tick().or(scope_ticker.tick()).await;
486 }
487 };
488// we don't care if it errors. If a scoped task errors it will propagate
489 // to get_results
490let _result = AssertUnwindSafe(executor.run(tick_forever))
491 .catch_unwind()
492 .await
493.is_ok();
494 }
495 };
496 get_results.or(execute_forever).await
497}
498499#[inline]
500async fn execute_external_scope<'scope, 'ticker, T>(
501 external_ticker: ThreadExecutorTicker<'scope, 'ticker>,
502 scope_ticker: ThreadExecutorTicker<'scope, 'ticker>,
503 get_results: impl Future<Output = Vec<T>>,
504 ) -> Vec<T> {
505let execute_forever = async {
506loop {
507let tick_forever = async {
508loop {
509 external_ticker.tick().or(scope_ticker.tick()).await;
510 }
511 };
512let _result = AssertUnwindSafe(tick_forever).catch_unwind().await.is_ok();
513 }
514 };
515 get_results.or(execute_forever).await
516}
517518#[inline]
519async fn execute_global_scope<'scope, 'ticker, T>(
520 executor: &'scope crate::executor::Executor<'scope>,
521 scope_ticker: ThreadExecutorTicker<'scope, 'ticker>,
522 get_results: impl Future<Output = Vec<T>>,
523 ) -> Vec<T> {
524let execute_forever = async {
525loop {
526let tick_forever = async {
527loop {
528 scope_ticker.tick().await;
529 }
530 };
531let _result = AssertUnwindSafe(executor.run(tick_forever))
532 .catch_unwind()
533 .await
534.is_ok();
535 }
536 };
537 get_results.or(execute_forever).await
538}
539540#[inline]
541async fn execute_scope<'scope, 'ticker, T>(
542 scope_ticker: ThreadExecutorTicker<'scope, 'ticker>,
543 get_results: impl Future<Output = Vec<T>>,
544 ) -> Vec<T> {
545let execute_forever = async {
546loop {
547let tick_forever = async {
548loop {
549 scope_ticker.tick().await;
550 }
551 };
552let _result = AssertUnwindSafe(tick_forever).catch_unwind().await.is_ok();
553 }
554 };
555 get_results.or(execute_forever).await
556}
557558/// Spawns a static future onto the thread pool. The returned [`Task`] is a
559 /// future that can be polled for the result. It can also be canceled and
560 /// "detached", allowing the task to continue running even if dropped. In
561 /// any case, the pool will execute the task even without polling by the
562 /// end-user.
563 ///
564 /// If the provided future is non-`Send`, [`TaskPool::spawn_local`] should
565 /// be used instead.
566pub fn spawn<T>(&self, future: impl Future<Output = T> + Send + 'static) -> Task<T>
567where
568T: Send + 'static,
569 {
570self.executor.spawn(future)
571 }
572573/// Spawns a static future on the thread-local async executor for the
574 /// current thread. The task will run entirely on the thread the task was
575 /// spawned on.
576 ///
577 /// The returned [`Task`] is a future that can be polled for the
578 /// result. It can also be canceled and "detached", allowing the task to
579 /// continue running even if dropped. In any case, the pool will execute the
580 /// task even without polling by the end-user.
581 ///
582 /// Users should generally prefer to use [`TaskPool::spawn`] instead,
583 /// unless the provided future is not `Send`.
584pub fn spawn_local<T>(&self, future: impl Future<Output = T> + 'static) -> Task<T>
585where
586T: 'static,
587 {
588TaskPool::LOCAL_EXECUTOR.with(|executor| executor.spawn(future))
589 }
590591/// Runs a function with the local executor. Typically used to tick
592 /// the local executor on the main thread as it needs to share time with
593 /// other things.
594 ///
595 /// ```
596 /// use bevy_tasks::TaskPool;
597 ///
598 /// TaskPool::new().with_local_executor(|local_executor| {
599 /// local_executor.try_tick();
600 /// });
601 /// ```
602pub fn with_local_executor<F, R>(&self, f: F) -> R
603where
604F: FnOnce(&crate::executor::LocalExecutor) -> R,
605 {
606Self::LOCAL_EXECUTOR.with(f)
607 }
608}
609610impl Defaultfor TaskPool {
611fn default() -> Self {
612Self::new()
613 }
614}
615616impl Dropfor TaskPool {
617fn drop(&mut self) {
618self.shutdown_tx.close();
619620let panicking = thread::panicking();
621for join_handle in self.threads.drain(..) {
622let res = join_handle.join();
623if !panicking {
624 res.expect("Task thread panicked while executing.");
625 }
626 }
627 }
628}
629630/// A [`TaskPool`] scope for running one or more non-`'static` futures.
631///
632/// For more information, see [`TaskPool::scope`].
633#[derive(#[automatically_derived]
impl<'scope, 'env: 'scope, T: ::core::fmt::Debug> ::core::fmt::Debug for
Scope<'scope, 'env, T> {
#[inline]
fn fmt(&self, f: &mut ::core::fmt::Formatter) -> ::core::fmt::Result {
let names: &'static _ =
&["executor", "external_executor", "scope_executor", "spawned",
"scope", "env"];
let values: &[&dyn ::core::fmt::Debug] =
&[&self.executor, &self.external_executor, &self.scope_executor,
&self.spawned, &self.scope, &&self.env];
::core::fmt::Formatter::debug_struct_fields_finish(f, "Scope", names,
values)
}
}Debug)]
634pub struct Scope<'scope, 'env: 'scope, T> {
635 executor: &'scope crate::executor::Executor<'scope>,
636 external_executor: &'scope ThreadExecutor<'scope>,
637 scope_executor: &'scope ThreadExecutor<'scope>,
638 spawned: &'scope ConcurrentQueue<FallibleTask<Result<T, Box<dyn core::any::Any + Send>>>>,
639// make `Scope` invariant over 'scope and 'env
640scope: PhantomData<&'scope mut &'scope ()>,
641 env: PhantomData<&'env mut &'env ()>,
642}
643644impl<'scope, 'env, T: Send + 'scope> Scope<'scope, 'env, T> {
645/// Spawns a scoped future onto the thread pool. The scope *must* outlive
646 /// the provided future. The results of the future will be returned as a part of
647 /// [`TaskPool::scope`]'s return value.
648 ///
649 /// For futures that should run on the thread `scope` is called on [`Scope::spawn_on_scope`] should be used
650 /// instead.
651 ///
652 /// For more information, see [`TaskPool::scope`].
653pub fn spawn<Fut: Future<Output = T> + 'scope + Send>(&self, f: Fut) {
654let task = self655 .executor
656 .spawn(AssertUnwindSafe(f).catch_unwind())
657 .fallible();
658// ConcurrentQueue only errors when closed or full, but we never
659 // close and use an unbounded queue, so it is safe to unwrap
660self.spawned.push(task).unwrap();
661 }
662663/// Spawns a scoped future onto the thread the scope is run on. The scope *must* outlive
664 /// the provided future. The results of the future will be returned as a part of
665 /// [`TaskPool::scope`]'s return value. Users should generally prefer to use
666 /// [`Scope::spawn`] instead, unless the provided future needs to run on the scope's thread.
667 ///
668 /// For more information, see [`TaskPool::scope`].
669pub fn spawn_on_scope<Fut: Future<Output = T> + 'scope + Send>(&self, f: Fut) {
670let task = self671 .scope_executor
672 .spawn(AssertUnwindSafe(f).catch_unwind())
673 .fallible();
674// ConcurrentQueue only errors when closed or full, but we never
675 // close and use an unbounded queue, so it is safe to unwrap
676self.spawned.push(task).unwrap();
677 }
678679/// Spawns a scoped future onto the thread of the external thread executor.
680 /// This is typically the main thread. The scope *must* outlive
681 /// the provided future. The results of the future will be returned as a part of
682 /// [`TaskPool::scope`]'s return value. Users should generally prefer to use
683 /// [`Scope::spawn`] instead, unless the provided future needs to run on the external thread.
684 ///
685 /// For more information, see [`TaskPool::scope`].
686pub fn spawn_on_external<Fut: Future<Output = T> + 'scope + Send>(&self, f: Fut) {
687let task = self688 .external_executor
689 .spawn(AssertUnwindSafe(f).catch_unwind())
690 .fallible();
691// ConcurrentQueue only errors when closed or full, but we never
692 // close and use an unbounded queue, so it is safe to unwrap
693self.spawned.push(task).unwrap();
694 }
695}
696697impl<'scope, 'env, T> Dropfor Scope<'scope, 'env, T>
698where
699T: 'scope,
700{
701fn drop(&mut self) {
702block_on(async {
703while let Ok(task) = self.spawned.pop() {
704 task.cancel().await;
705 }
706 });
707 }
708}
709710#[cfg(test)]
711mod tests {
712use super::*;
713use core::sync::atomic::{AtomicBool, AtomicI32, Ordering};
714use std::sync::Barrier;
715716#[test]
717fn test_spawn() {
718let pool = TaskPool::new();
719720let foo = Box::new(42);
721let foo = &*foo;
722723let count = Arc::new(AtomicI32::new(0));
724725let outputs = pool.scope(|scope| {
726for _ in 0..100 {
727let count_clone = count.clone();
728 scope.spawn(async move {
729if *foo != 42 {
730panic!("not 42!?!?")
731 } else {
732 count_clone.fetch_add(1, Ordering::Relaxed);
733*foo
734 }
735 });
736 }
737 });
738739for output in &outputs {
740assert_eq!(*output, 42);
741 }
742743assert_eq!(outputs.len(), 100);
744assert_eq!(count.load(Ordering::Relaxed), 100);
745 }
746747#[test]
748fn test_thread_callbacks() {
749let counter = Arc::new(AtomicI32::new(0));
750let start_counter = counter.clone();
751 {
752let barrier = Arc::new(Barrier::new(11));
753let last_barrier = barrier.clone();
754// Build and immediately drop to terminate
755let _pool = TaskPoolBuilder::new()
756 .num_threads(10)
757 .on_thread_spawn(move || {
758 start_counter.fetch_add(1, Ordering::Relaxed);
759 barrier.clone().wait();
760 })
761 .build();
762 last_barrier.wait();
763assert_eq!(10, counter.load(Ordering::Relaxed));
764 }
765assert_eq!(10, counter.load(Ordering::Relaxed));
766let end_counter = counter.clone();
767 {
768let _pool = TaskPoolBuilder::new()
769 .num_threads(20)
770 .on_thread_destroy(move || {
771 end_counter.fetch_sub(1, Ordering::Relaxed);
772 })
773 .build();
774assert_eq!(10, counter.load(Ordering::Relaxed));
775 }
776assert_eq!(-10, counter.load(Ordering::Relaxed));
777let start_counter = counter.clone();
778let end_counter = counter.clone();
779 {
780let barrier = Arc::new(Barrier::new(6));
781let last_barrier = barrier.clone();
782let _pool = TaskPoolBuilder::new()
783 .num_threads(5)
784 .on_thread_spawn(move || {
785 start_counter.fetch_add(1, Ordering::Relaxed);
786 barrier.wait();
787 })
788 .on_thread_destroy(move || {
789 end_counter.fetch_sub(1, Ordering::Relaxed);
790 })
791 .build();
792 last_barrier.wait();
793assert_eq!(-5, counter.load(Ordering::Relaxed));
794 }
795assert_eq!(-10, counter.load(Ordering::Relaxed));
796 }
797798#[test]
799fn test_mixed_spawn_on_scope_and_spawn() {
800let pool = TaskPool::new();
801802let foo = Box::new(42);
803let foo = &*foo;
804805let local_count = Arc::new(AtomicI32::new(0));
806let non_local_count = Arc::new(AtomicI32::new(0));
807808let outputs = pool.scope(|scope| {
809for i in 0..100 {
810if i % 2 == 0 {
811let count_clone = non_local_count.clone();
812 scope.spawn(async move {
813if *foo != 42 {
814panic!("not 42!?!?")
815 } else {
816 count_clone.fetch_add(1, Ordering::Relaxed);
817*foo
818 }
819 });
820 } else {
821let count_clone = local_count.clone();
822 scope.spawn_on_scope(async move {
823if *foo != 42 {
824panic!("not 42!?!?")
825 } else {
826 count_clone.fetch_add(1, Ordering::Relaxed);
827*foo
828 }
829 });
830 }
831 }
832 });
833834for output in &outputs {
835assert_eq!(*output, 42);
836 }
837838assert_eq!(outputs.len(), 100);
839assert_eq!(local_count.load(Ordering::Relaxed), 50);
840assert_eq!(non_local_count.load(Ordering::Relaxed), 50);
841 }
842843#[test]
844fn test_thread_locality() {
845let pool = Arc::new(TaskPool::new());
846let count = Arc::new(AtomicI32::new(0));
847let barrier = Arc::new(Barrier::new(101));
848let thread_check_failed = Arc::new(AtomicBool::new(false));
849850for _ in 0..100 {
851let inner_barrier = barrier.clone();
852let count_clone = count.clone();
853let inner_pool = pool.clone();
854let inner_thread_check_failed = thread_check_failed.clone();
855 thread::spawn(move || {
856 inner_pool.scope(|scope| {
857let inner_count_clone = count_clone.clone();
858 scope.spawn(async move {
859 inner_count_clone.fetch_add(1, Ordering::Release);
860 });
861let spawner = thread::current().id();
862let inner_count_clone = count_clone.clone();
863 scope.spawn_on_scope(async move {
864 inner_count_clone.fetch_add(1, Ordering::Release);
865if thread::current().id() != spawner {
866// NOTE: This check is using an atomic rather than simply panicking the
867 // thread to avoid deadlocking the barrier on failure
868inner_thread_check_failed.store(true, Ordering::Release);
869 }
870 });
871 });
872 inner_barrier.wait();
873 });
874 }
875 barrier.wait();
876assert!(!thread_check_failed.load(Ordering::Acquire));
877assert_eq!(count.load(Ordering::Acquire), 200);
878 }
879880#[test]
881fn test_nested_spawn() {
882let pool = TaskPool::new();
883884let foo = Box::new(42);
885let foo = &*foo;
886887let count = Arc::new(AtomicI32::new(0));
888889let outputs: Vec<i32> = pool.scope(|scope| {
890for _ in 0..10 {
891let count_clone = count.clone();
892 scope.spawn(async move {
893for _ in 0..10 {
894let count_clone_clone = count_clone.clone();
895 scope.spawn(async move {
896if *foo != 42 {
897panic!("not 42!?!?")
898 } else {
899 count_clone_clone.fetch_add(1, Ordering::Relaxed);
900*foo
901 }
902 });
903 }
904*foo
905 });
906 }
907 });
908909for output in &outputs {
910assert_eq!(*output, 42);
911 }
912913// the inner loop runs 100 times and the outer one runs 10. 100 + 10
914assert_eq!(outputs.len(), 110);
915assert_eq!(count.load(Ordering::Relaxed), 100);
916 }
917918#[test]
919fn test_nested_locality() {
920let pool = Arc::new(TaskPool::new());
921let count = Arc::new(AtomicI32::new(0));
922let barrier = Arc::new(Barrier::new(101));
923let thread_check_failed = Arc::new(AtomicBool::new(false));
924925for _ in 0..100 {
926let inner_barrier = barrier.clone();
927let count_clone = count.clone();
928let inner_pool = pool.clone();
929let inner_thread_check_failed = thread_check_failed.clone();
930 thread::spawn(move || {
931 inner_pool.scope(|scope| {
932let spawner = thread::current().id();
933let inner_count_clone = count_clone.clone();
934 scope.spawn(async move {
935 inner_count_clone.fetch_add(1, Ordering::Release);
936937// spawning on the scope from another thread runs the futures on the scope's thread
938scope.spawn_on_scope(async move {
939 inner_count_clone.fetch_add(1, Ordering::Release);
940if thread::current().id() != spawner {
941// NOTE: This check is using an atomic rather than simply panicking the
942 // thread to avoid deadlocking the barrier on failure
943inner_thread_check_failed.store(true, Ordering::Release);
944 }
945 });
946 });
947 });
948 inner_barrier.wait();
949 });
950 }
951 barrier.wait();
952assert!(!thread_check_failed.load(Ordering::Acquire));
953assert_eq!(count.load(Ordering::Acquire), 200);
954 }
955956// This test will often freeze on other executors.
957#[test]
958fn test_nested_scopes() {
959let pool = TaskPool::new();
960let count = Arc::new(AtomicI32::new(0));
961962 pool.scope(|scope| {
963 scope.spawn(async {
964 pool.scope(|scope| {
965 scope.spawn(async {
966 count.fetch_add(1, Ordering::Relaxed);
967 });
968 });
969 });
970 });
971972assert_eq!(count.load(Ordering::Acquire), 1);
973 }
974}