Skip to main content

bevy_tasks/
task_pool.rs

1use alloc::{boxed::Box, format, string::String, vec::Vec};
2use core::{future::Future, marker::PhantomData, mem, panic::AssertUnwindSafe};
3use std::{
4    thread::{self, JoinHandle},
5    thread_local,
6};
7
8use crate::executor::FallibleTask;
9use bevy_platform::sync::Arc;
10use concurrent_queue::ConcurrentQueue;
11use futures_lite::FutureExt;
12
13use crate::{
14    block_on,
15    thread_executor::{ThreadExecutor, ThreadExecutorTicker},
16    Task,
17};
18
19struct CallOnDrop(Option<Arc<dyn Fn() + Send + Sync + 'static>>);
20
21impl Drop for CallOnDrop {
22    fn drop(&mut self) {
23        if let Some(call) = self.0.as_ref() {
24            call();
25        }
26    }
27}
28
29/// 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
35    num_threads: Option<usize>,
36    /// If set, we'll use the given stack size rather than the system default
37    stack_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)"`.
40    thread_name: Option<String>,
41
42    on_thread_spawn: Option<Arc<dyn Fn() + Send + Sync + 'static>>,
43    on_thread_destroy: Option<Arc<dyn Fn() + Send + Sync + 'static>>,
44}
45
46impl TaskPoolBuilder {
47    /// Creates a new [`TaskPoolBuilder`] instance
48    pub fn new() -> Self {
49        Self::default()
50    }
51
52    /// Override the number of threads created for the pool. If unset, we default to the number
53    /// of logical cores of the system
54    pub fn num_threads(mut self, num_threads: usize) -> Self {
55        self.num_threads = Some(num_threads);
56        self
57    }
58
59    /// Override the stack size of the threads created for the pool
60    pub fn stack_size(mut self, stack_size: usize) -> Self {
61        self.stack_size = Some(stack_size);
62        self
63    }
64
65    /// 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)`
67    pub fn thread_name(mut self, thread_name: String) -> Self {
68        self.thread_name = Some(thread_name);
69        self
70    }
71
72    /// 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.
76    pub fn on_thread_spawn(mut self, f: impl Fn() + Send + Sync + 'static) -> Self {
77        let arc = Arc::new(f);
78
79        #[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`
87        let arc = unsafe {
88            Arc::from_raw(Arc::into_raw(arc) as *const (dyn Fn() + Send + Sync + 'static))
89        };
90
91        self.on_thread_spawn = Some(arc);
92        self
93    }
94
95    /// 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.
99    pub fn on_thread_destroy(mut self, f: impl Fn() + Send + Sync + 'static) -> Self {
100        let arc = Arc::new(f);
101
102        #[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`
110        let arc = unsafe {
111            Arc::from_raw(Arc::into_raw(arc) as *const (dyn Fn() + Send + Sync + 'static))
112        };
113
114        self.on_thread_destroy = Some(arc);
115        self
116    }
117
118    /// Creates a new [`TaskPool`] based on the current options.
119    pub fn build(self) -> TaskPool {
120        TaskPool::new_internal(self)
121    }
122}
123
124/// 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.
137    executor: Arc<crate::executor::Executor<'static>>,
138
139    // The inner state of the pool.
140    threads: Vec<JoinHandle<()>>,
141    shutdown_tx: async_channel::Sender<()>,
142}
143
144impl TaskPool {
145    const 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! {
146        static LOCAL_EXECUTOR: crate::executor::LocalExecutor<'static> = const { crate::executor::LocalExecutor::new() };
147        static THREAD_EXECUTOR: Arc<ThreadExecutor<'static>> = Arc::new(ThreadExecutor::new());
148    }
149
150    /// Each thread should only create one `ThreadExecutor`, otherwise, there are good chances they will deadlock
151    pub fn get_thread_executor() -> Arc<ThreadExecutor<'static>> {
152        Self::THREAD_EXECUTOR.with(Clone::clone)
153    }
154
155    /// Create a `TaskPool` with the default configuration.
156    pub fn new() -> Self {
157        TaskPoolBuilder::new().build()
158    }
159
160    fn new_internal(builder: TaskPoolBuilder) -> Self {
161        let (shutdown_tx, shutdown_rx) = async_channel::unbounded::<()>();
162
163        let executor = Arc::new(crate::executor::Executor::new());
164
165        let num_requested_threads = builder
166            .num_threads
167            .unwrap_or_else(crate::available_parallelism);
168
169        // 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"))]
172        let num_threads = num_requested_threads.max(2);
173        #[cfg(not(all(test, feature = "multi_threaded")))]
174        let num_threads = num_requested_threads;
175
176        let threads = (0..num_threads)
177            .map(|i| {
178                let ex = Arc::clone(&executor);
179                let shutdown_rx = shutdown_rx.clone();
180
181                let 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                };
186                let mut thread_builder = thread::Builder::new().name(thread_name);
187
188                if let Some(stack_size) = builder.stack_size {
189                    thread_builder = thread_builder.stack_size(stack_size);
190                }
191
192                let on_thread_spawn = builder.on_thread_spawn.clone();
193                let on_thread_destroy = builder.on_thread_destroy.clone();
194
195                thread_builder
196                    .spawn(move || {
197                        TaskPool::LOCAL_EXECUTOR.with(|local_executor| {
198                            if let Some(on_thread_spawn) = on_thread_spawn {
199                                on_thread_spawn();
200                                drop(on_thread_spawn);
201                            }
202                            let _destructor = CallOnDrop(on_thread_destroy);
203                            loop {
204                                let res = std::panic::catch_unwind(|| {
205                                    let tick_forever = async move {
206                                        loop {
207                                            local_executor.tick().await;
208                                        }
209                                    };
210                                    block_on(ex.run(tick_forever.or(shutdown_rx.recv())))
211                                });
212                                if let Ok(value) = res {
213                                    // Use unwrap_err because we expect a Closed error
214                                    value.unwrap_err();
215                                    break;
216                                }
217                            }
218                        });
219                    })
220                    .expect("Failed to spawn thread.")
221            })
222            .collect();
223
224        Self {
225            executor,
226            threads,
227            shutdown_tx,
228        }
229    }
230
231    /// Return the number of threads owned by the task pool
232    pub fn thread_num(&self) -> usize {
233        self.threads.len()
234    }
235
236    /// 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    /// }
317    pub fn scope<'env, F, T>(&self, f: F) -> Vec<T>
318    where
319        F: for<'scope> FnOnce(&'scope Scope<'scope, 'env, T>),
320        T: Send + 'static,
321    {
322        Self::THREAD_EXECUTOR.with(|scope_executor| {
323            self.scope_with_executor_inner(true, scope_executor, scope_executor, f)
324        })
325    }
326
327    /// 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.
336    pub 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>
342    where
343        F: for<'scope> FnOnce(&'scope Scope<'scope, 'env, T>),
344        T: Send + 'static,
345    {
346        Self::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.
349            if let Some(external_executor) = external_executor {
350                self.scope_with_executor_inner(
351                    tick_task_pool_executor,
352                    external_executor,
353                    scope_executor,
354                    f,
355                )
356            } else {
357                self.scope_with_executor_inner(
358                    tick_task_pool_executor,
359                    scope_executor,
360                    scope_executor,
361                    f,
362                )
363            }
364        })
365    }
366
367    #[expect(unsafe_code, reason = "Required to transmute lifetimes.")]
368    fn 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>
375    where
376        F: 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.
386        let executor: &crate::executor::Executor = &self.executor;
387        // SAFETY: As above, all futures must complete in this function so we can change the lifetime
388        let 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
390        let external_executor: &'env ThreadExecutor<'env> =
391            unsafe { mem::transmute(external_executor) };
392        // SAFETY: As above, all futures must complete in this function so we can change the lifetime
393        let scope_executor: &'env ThreadExecutor<'env> = unsafe { mem::transmute(scope_executor) };
394        let spawned: ConcurrentQueue<FallibleTask<Result<T, Box<dyn core::any::Any + Send>>>> =
395            ConcurrentQueue::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
398        let spawned: &'env ConcurrentQueue<
399            FallibleTask<Result<T, Box<dyn core::any::Any + Send>>>,
400        > = unsafe { mem::transmute(&spawned) };
401
402        let scope = Scope {
403            executor,
404            external_executor,
405            scope_executor,
406            spawned,
407            scope: PhantomData,
408            env: PhantomData,
409        };
410
411        // 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
413        let scope: &'env Scope<'_, 'env, T> = unsafe { mem::transmute(&scope) };
414
415        f(scope);
416
417        if spawned.is_empty() {
418            Vec::new()
419        } else {
420            block_on(async move {
421                let get_results = async {
422                    let mut results = Vec::with_capacity(spawned.len());
423                    while let Ok(task) = spawned.pop() {
424                        if let Some(res) = task.await {
425                            match res {
426                                Ok(res) => results.push(res),
427                                Err(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                    }
433                    results
434                };
435
436                let tick_task_pool_executor = tick_task_pool_executor || self.threads.is_empty();
437
438                // 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.
441                let scope_ticker = scope_executor.ticker().unwrap();
442                let external_ticker = if !external_executor.is_same(scope_executor) {
443                    external_executor.ticker()
444                } else {
445                    None
446                };
447
448                match (external_ticker, tick_task_pool_executor) {
449                    (Some(external_ticker), true) => {
450                        Self::execute_global_external_scope(
451                            executor,
452                            external_ticker,
453                            scope_ticker,
454                            get_results,
455                        )
456                        .await
457                    }
458                    (Some(external_ticker), false) => {
459                        Self::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) => {
464                        Self::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    }
471
472    #[inline]
473    async 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
481        let execute_forever = async move {
482            loop {
483                let tick_forever = async {
484                    loop {
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
490                let _result = AssertUnwindSafe(executor.run(tick_forever))
491                    .catch_unwind()
492                    .await
493                    .is_ok();
494            }
495        };
496        get_results.or(execute_forever).await
497    }
498
499    #[inline]
500    async 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> {
505        let execute_forever = async {
506            loop {
507                let tick_forever = async {
508                    loop {
509                        external_ticker.tick().or(scope_ticker.tick()).await;
510                    }
511                };
512                let _result = AssertUnwindSafe(tick_forever).catch_unwind().await.is_ok();
513            }
514        };
515        get_results.or(execute_forever).await
516    }
517
518    #[inline]
519    async 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> {
524        let execute_forever = async {
525            loop {
526                let tick_forever = async {
527                    loop {
528                        scope_ticker.tick().await;
529                    }
530                };
531                let _result = AssertUnwindSafe(executor.run(tick_forever))
532                    .catch_unwind()
533                    .await
534                    .is_ok();
535            }
536        };
537        get_results.or(execute_forever).await
538    }
539
540    #[inline]
541    async fn execute_scope<'scope, 'ticker, T>(
542        scope_ticker: ThreadExecutorTicker<'scope, 'ticker>,
543        get_results: impl Future<Output = Vec<T>>,
544    ) -> Vec<T> {
545        let execute_forever = async {
546            loop {
547                let tick_forever = async {
548                    loop {
549                        scope_ticker.tick().await;
550                    }
551                };
552                let _result = AssertUnwindSafe(tick_forever).catch_unwind().await.is_ok();
553            }
554        };
555        get_results.or(execute_forever).await
556    }
557
558    /// 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.
566    pub fn spawn<T>(&self, future: impl Future<Output = T> + Send + 'static) -> Task<T>
567    where
568        T: Send + 'static,
569    {
570        self.executor.spawn(future)
571    }
572
573    /// 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`.
584    pub fn spawn_local<T>(&self, future: impl Future<Output = T> + 'static) -> Task<T>
585    where
586        T: 'static,
587    {
588        TaskPool::LOCAL_EXECUTOR.with(|executor| executor.spawn(future))
589    }
590
591    /// 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    /// ```
602    pub fn with_local_executor<F, R>(&self, f: F) -> R
603    where
604        F: FnOnce(&crate::executor::LocalExecutor) -> R,
605    {
606        Self::LOCAL_EXECUTOR.with(f)
607    }
608}
609
610impl Default for TaskPool {
611    fn default() -> Self {
612        Self::new()
613    }
614}
615
616impl Drop for TaskPool {
617    fn drop(&mut self) {
618        self.shutdown_tx.close();
619
620        let panicking = thread::panicking();
621        for join_handle in self.threads.drain(..) {
622            let res = join_handle.join();
623            if !panicking {
624                res.expect("Task thread panicked while executing.");
625            }
626        }
627    }
628}
629
630/// 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
640    scope: PhantomData<&'scope mut &'scope ()>,
641    env: PhantomData<&'env mut &'env ()>,
642}
643
644impl<'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`].
653    pub fn spawn<Fut: Future<Output = T> + 'scope + Send>(&self, f: Fut) {
654        let task = self
655            .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
660        self.spawned.push(task).unwrap();
661    }
662
663    /// 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`].
669    pub fn spawn_on_scope<Fut: Future<Output = T> + 'scope + Send>(&self, f: Fut) {
670        let task = self
671            .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
676        self.spawned.push(task).unwrap();
677    }
678
679    /// 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`].
686    pub fn spawn_on_external<Fut: Future<Output = T> + 'scope + Send>(&self, f: Fut) {
687        let task = self
688            .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
693        self.spawned.push(task).unwrap();
694    }
695}
696
697impl<'scope, 'env, T> Drop for Scope<'scope, 'env, T>
698where
699    T: 'scope,
700{
701    fn drop(&mut self) {
702        block_on(async {
703            while let Ok(task) = self.spawned.pop() {
704                task.cancel().await;
705            }
706        });
707    }
708}
709
710#[cfg(test)]
711mod tests {
712    use super::*;
713    use core::sync::atomic::{AtomicBool, AtomicI32, Ordering};
714    use std::sync::Barrier;
715
716    #[test]
717    fn test_spawn() {
718        let pool = TaskPool::new();
719
720        let foo = Box::new(42);
721        let foo = &*foo;
722
723        let count = Arc::new(AtomicI32::new(0));
724
725        let outputs = pool.scope(|scope| {
726            for _ in 0..100 {
727                let count_clone = count.clone();
728                scope.spawn(async move {
729                    if *foo != 42 {
730                        panic!("not 42!?!?")
731                    } else {
732                        count_clone.fetch_add(1, Ordering::Relaxed);
733                        *foo
734                    }
735                });
736            }
737        });
738
739        for output in &outputs {
740            assert_eq!(*output, 42);
741        }
742
743        assert_eq!(outputs.len(), 100);
744        assert_eq!(count.load(Ordering::Relaxed), 100);
745    }
746
747    #[test]
748    fn test_thread_callbacks() {
749        let counter = Arc::new(AtomicI32::new(0));
750        let start_counter = counter.clone();
751        {
752            let barrier = Arc::new(Barrier::new(11));
753            let last_barrier = barrier.clone();
754            // Build and immediately drop to terminate
755            let _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();
763            assert_eq!(10, counter.load(Ordering::Relaxed));
764        }
765        assert_eq!(10, counter.load(Ordering::Relaxed));
766        let end_counter = counter.clone();
767        {
768            let _pool = TaskPoolBuilder::new()
769                .num_threads(20)
770                .on_thread_destroy(move || {
771                    end_counter.fetch_sub(1, Ordering::Relaxed);
772                })
773                .build();
774            assert_eq!(10, counter.load(Ordering::Relaxed));
775        }
776        assert_eq!(-10, counter.load(Ordering::Relaxed));
777        let start_counter = counter.clone();
778        let end_counter = counter.clone();
779        {
780            let barrier = Arc::new(Barrier::new(6));
781            let last_barrier = barrier.clone();
782            let _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();
793            assert_eq!(-5, counter.load(Ordering::Relaxed));
794        }
795        assert_eq!(-10, counter.load(Ordering::Relaxed));
796    }
797
798    #[test]
799    fn test_mixed_spawn_on_scope_and_spawn() {
800        let pool = TaskPool::new();
801
802        let foo = Box::new(42);
803        let foo = &*foo;
804
805        let local_count = Arc::new(AtomicI32::new(0));
806        let non_local_count = Arc::new(AtomicI32::new(0));
807
808        let outputs = pool.scope(|scope| {
809            for i in 0..100 {
810                if i % 2 == 0 {
811                    let count_clone = non_local_count.clone();
812                    scope.spawn(async move {
813                        if *foo != 42 {
814                            panic!("not 42!?!?")
815                        } else {
816                            count_clone.fetch_add(1, Ordering::Relaxed);
817                            *foo
818                        }
819                    });
820                } else {
821                    let count_clone = local_count.clone();
822                    scope.spawn_on_scope(async move {
823                        if *foo != 42 {
824                            panic!("not 42!?!?")
825                        } else {
826                            count_clone.fetch_add(1, Ordering::Relaxed);
827                            *foo
828                        }
829                    });
830                }
831            }
832        });
833
834        for output in &outputs {
835            assert_eq!(*output, 42);
836        }
837
838        assert_eq!(outputs.len(), 100);
839        assert_eq!(local_count.load(Ordering::Relaxed), 50);
840        assert_eq!(non_local_count.load(Ordering::Relaxed), 50);
841    }
842
843    #[test]
844    fn test_thread_locality() {
845        let pool = Arc::new(TaskPool::new());
846        let count = Arc::new(AtomicI32::new(0));
847        let barrier = Arc::new(Barrier::new(101));
848        let thread_check_failed = Arc::new(AtomicBool::new(false));
849
850        for _ in 0..100 {
851            let inner_barrier = barrier.clone();
852            let count_clone = count.clone();
853            let inner_pool = pool.clone();
854            let inner_thread_check_failed = thread_check_failed.clone();
855            thread::spawn(move || {
856                inner_pool.scope(|scope| {
857                    let inner_count_clone = count_clone.clone();
858                    scope.spawn(async move {
859                        inner_count_clone.fetch_add(1, Ordering::Release);
860                    });
861                    let spawner = thread::current().id();
862                    let inner_count_clone = count_clone.clone();
863                    scope.spawn_on_scope(async move {
864                        inner_count_clone.fetch_add(1, Ordering::Release);
865                        if 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
868                            inner_thread_check_failed.store(true, Ordering::Release);
869                        }
870                    });
871                });
872                inner_barrier.wait();
873            });
874        }
875        barrier.wait();
876        assert!(!thread_check_failed.load(Ordering::Acquire));
877        assert_eq!(count.load(Ordering::Acquire), 200);
878    }
879
880    #[test]
881    fn test_nested_spawn() {
882        let pool = TaskPool::new();
883
884        let foo = Box::new(42);
885        let foo = &*foo;
886
887        let count = Arc::new(AtomicI32::new(0));
888
889        let outputs: Vec<i32> = pool.scope(|scope| {
890            for _ in 0..10 {
891                let count_clone = count.clone();
892                scope.spawn(async move {
893                    for _ in 0..10 {
894                        let count_clone_clone = count_clone.clone();
895                        scope.spawn(async move {
896                            if *foo != 42 {
897                                panic!("not 42!?!?")
898                            } else {
899                                count_clone_clone.fetch_add(1, Ordering::Relaxed);
900                                *foo
901                            }
902                        });
903                    }
904                    *foo
905                });
906            }
907        });
908
909        for output in &outputs {
910            assert_eq!(*output, 42);
911        }
912
913        // the inner loop runs 100 times and the outer one runs 10. 100 + 10
914        assert_eq!(outputs.len(), 110);
915        assert_eq!(count.load(Ordering::Relaxed), 100);
916    }
917
918    #[test]
919    fn test_nested_locality() {
920        let pool = Arc::new(TaskPool::new());
921        let count = Arc::new(AtomicI32::new(0));
922        let barrier = Arc::new(Barrier::new(101));
923        let thread_check_failed = Arc::new(AtomicBool::new(false));
924
925        for _ in 0..100 {
926            let inner_barrier = barrier.clone();
927            let count_clone = count.clone();
928            let inner_pool = pool.clone();
929            let inner_thread_check_failed = thread_check_failed.clone();
930            thread::spawn(move || {
931                inner_pool.scope(|scope| {
932                    let spawner = thread::current().id();
933                    let inner_count_clone = count_clone.clone();
934                    scope.spawn(async move {
935                        inner_count_clone.fetch_add(1, Ordering::Release);
936
937                        // spawning on the scope from another thread runs the futures on the scope's thread
938                        scope.spawn_on_scope(async move {
939                            inner_count_clone.fetch_add(1, Ordering::Release);
940                            if 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
943                                inner_thread_check_failed.store(true, Ordering::Release);
944                            }
945                        });
946                    });
947                });
948                inner_barrier.wait();
949            });
950        }
951        barrier.wait();
952        assert!(!thread_check_failed.load(Ordering::Acquire));
953        assert_eq!(count.load(Ordering::Acquire), 200);
954    }
955
956    // This test will often freeze on other executors.
957    #[test]
958    fn test_nested_scopes() {
959        let pool = TaskPool::new();
960        let count = Arc::new(AtomicI32::new(0));
961
962        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        });
971
972        assert_eq!(count.load(Ordering::Acquire), 1);
973    }
974}