1use std::{
12 any::Any,
13 marker::PhantomData,
14 num::NonZeroUsize,
15 panic::{self, AssertUnwindSafe, catch_unwind},
16 sync::{
17 Arc,
18 atomic::{AtomicUsize, Ordering},
19 mpmc::{self, Receiver, Sender},
20 },
21 thread::{self, Thread},
22 time::{Duration, Instant},
23};
24
25use parking_lot::Mutex;
26use tokio::{runtime::Handle, task::block_in_place};
27use tracing::{Span, info_span};
28
29use crate::{TurboTasksApi, manager::try_turbo_tasks, turbo_tasks_scope};
30
31type WorkQueueJob = (usize, Box<dyn FnOnce() + Send + 'static>);
33
34struct ScopeInner {
35 main_thread: Thread,
36 remaining_tasks: AtomicUsize,
37 panic: Mutex<Option<(Box<dyn Any + Send + 'static>, usize)>>,
40 work_queue: Receiver<WorkQueueJob>,
43}
44
45impl ScopeInner {
46 fn on_task_finished(&self, panic: Option<(Box<dyn Any + Send + 'static>, usize)>) {
47 if let Some((err, index)) = panic {
48 let mut old_panic = self.panic.lock();
49 if old_panic.as_ref().is_none_or(|&(_, i)| i > index) {
50 *old_panic = Some((err, index));
51 }
52 }
53 if self.remaining_tasks.fetch_sub(1, Ordering::Release) == 1 {
54 self.main_thread.unpark();
55 }
56 }
57
58 fn wait(&self) {
59 if self.remaining_tasks.load(Ordering::Acquire) == 0 {
60 return;
61 }
62
63 let _span = info_span!("blocking").entered();
64
65 const TIMEOUT: Duration = Duration::from_millis(1);
67 let beginning_park = Instant::now();
68
69 let mut timeout_remaining = TIMEOUT;
70 loop {
71 thread::park_timeout(timeout_remaining);
72 if self.remaining_tasks.load(Ordering::Acquire) == 0 {
73 return;
74 }
75 let elapsed = beginning_park.elapsed();
76 if elapsed >= TIMEOUT {
77 break;
78 }
79 timeout_remaining = TIMEOUT - elapsed;
80 }
81
82 block_in_place(|| {
84 while self.remaining_tasks.load(Ordering::Acquire) != 0 {
85 thread::park();
86 }
87 });
88 }
89
90 fn wait_and_rethrow_panic(&self) {
91 self.wait();
92 if let Some((err, _)) = self.panic.lock().take() {
93 panic::resume_unwind(err);
94 }
95 }
96
97 fn run_jobs(&self) {
101 while let Ok((index, job)) = self.work_queue.recv() {
102 let result = catch_unwind(AssertUnwindSafe(job));
103 let panic = result.err().map(|e| (e, index));
104 self.on_task_finished(panic);
105 }
106 }
107}
108
109pub struct Scope<'scope, 'env: 'scope, R: Send + 'env> {
113 results: &'scope [Mutex<Option<R>>],
114 index: AtomicUsize,
115 inner: Arc<ScopeInner>,
116 work_queue: Option<Sender<WorkQueueJob>>,
118 handle: Handle,
119 worker_tasks: NonZeroUsize,
122 turbo_tasks: Option<Arc<dyn TurboTasksApi>>,
123 span: Span,
124 env: PhantomData<&'env mut &'env ()>,
129}
130
131impl<'scope, 'env: 'scope, R: Send + 'env> Scope<'scope, 'env, R> {
132 unsafe fn new(results: &'scope [Mutex<Option<R>>]) -> Self {
138 let handle = Handle::current();
139 let worker_tasks = NonZeroUsize::new(handle.metrics().num_workers().min(results.len()))
141 .unwrap_or(NonZeroUsize::MIN);
142 let (sender, receiver) = mpmc::channel();
143 Self {
144 results,
145 index: AtomicUsize::new(0),
146 inner: Arc::new(ScopeInner {
147 main_thread: thread::current(),
148 remaining_tasks: AtomicUsize::new(0),
149 panic: Mutex::new(None),
150 work_queue: receiver,
151 }),
152 work_queue: Some(sender),
153 handle,
154 worker_tasks,
155 turbo_tasks: try_turbo_tasks(),
156 span: Span::current(),
157 env: PhantomData,
158 }
159 }
160
161 pub fn spawn<F>(&self, f: F)
163 where
164 F: FnOnce() -> R + Send + 'env,
165 {
166 let index = self.index.fetch_add(1, Ordering::Relaxed);
167 assert!(index < self.results.len(), "Too many tasks spawned");
168 let result_cell: &Mutex<Option<R>> = &self.results[index];
169
170 let turbo_tasks = self.turbo_tasks.clone();
171 let f: Box<dyn FnOnce() + Send + 'scope> = Box::new(|| {
172 let result = {
173 if let Some(turbo_tasks) = turbo_tasks {
174 turbo_tasks_scope(turbo_tasks, f)
176 } else {
177 f()
179 }
180 };
181 *result_cell.lock() = Some(result);
182 });
183 let f: *mut (dyn FnOnce() + Send + 'scope) = Box::into_raw(f);
184
185 let f = unsafe {
188 std::mem::transmute::<
189 *mut (dyn FnOnce() + Send + 'scope),
190 *mut (dyn FnOnce() + Send + 'static),
191 >(f)
192 };
193
194 let f = unsafe { Box::from_raw(f) };
196
197 self.inner.remaining_tasks.fetch_add(1, Ordering::Relaxed);
198
199 self.work_queue
203 .as_ref()
204 .expect("sender is only taken in Drop")
205 .send((index, f))
206 .expect("receiver is owned by inner and outlives the scope");
207
208 if index < self.worker_tasks.get() - 1 {
210 let inner = self.inner.clone();
211 let span = self.span.clone();
212 self.handle.spawn(async move {
213 let _span = span.entered();
214 inner.run_jobs();
215 });
216 }
217 }
218}
219
220impl<'scope, 'env: 'scope, R: Send + 'env> Drop for Scope<'scope, 'env, R> {
221 fn drop(&mut self) {
222 drop(
226 self.work_queue
227 .take()
228 .expect("sender is taken exactly once, here in Drop"),
229 );
230 self.inner.run_jobs();
232 self.inner.wait_and_rethrow_panic();
233 }
234}
235
236pub fn scope_bounded<'env, F, R>(number_of_tasks: usize, f: F) -> impl Iterator<Item = R>
251where
252 R: Send + 'env,
253 F: for<'scope> FnOnce(&'scope Scope<'scope, 'env, R>) + 'env,
254{
255 let mut results = Vec::with_capacity(number_of_tasks);
256 for _ in 0..number_of_tasks {
257 results.push(Mutex::new(None));
258 }
259 let results = results.into_boxed_slice();
260 let result = {
261 let scope = unsafe { Scope::new(&results) };
263 catch_unwind(AssertUnwindSafe(|| f(&scope)))
264 };
265 if let Err(panic) = result {
266 panic::resume_unwind(panic);
267 }
268 results.into_iter().map(|mutex| {
269 mutex
270 .into_inner()
271 .expect("All values are set when the scope returns without panic")
272 })
273}
274
275#[cfg(test)]
276mod tests {
277 use std::{
278 panic::{AssertUnwindSafe, catch_unwind},
279 sync::atomic::AtomicUsize,
280 };
281
282 use super::*;
283
284 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
291 async fn test_scope_worker_threads_occupied() {
292 const WORKER_THREADS: usize = 2;
293 const JOBS: usize = 64;
294 const RELEASE_AFTER: Duration = Duration::from_secs(4);
295
296 let ready = Arc::new(AtomicUsize::new(0));
299 let mut occupiers = Vec::with_capacity(WORKER_THREADS);
300 for _ in 0..WORKER_THREADS {
301 let ready = ready.clone();
302 occupiers.push(tokio::spawn(async move {
303 ready.fetch_add(1, Ordering::SeqCst);
304 thread::sleep(RELEASE_AFTER);
305 }));
306 }
307 while ready.load(Ordering::SeqCst) < WORKER_THREADS {
309 tokio::task::yield_now().await;
310 }
311
312 let started = Instant::now();
313 let results = tokio::task::spawn_blocking(move || {
314 scope_bounded(JOBS, |scope| {
315 for i in 0..JOBS {
316 scope.spawn(move || i);
317 }
318 })
319 .collect::<Vec<_>>()
320 })
321 .await
322 .unwrap();
323 let elapsed = started.elapsed();
324
325 assert_eq!(results.len(), JOBS);
326 results.iter().enumerate().for_each(|(i, &result)| {
327 assert_eq!(result, i);
328 });
329 assert!(
330 elapsed < RELEASE_AFTER / 2,
331 "scope_bounded took {elapsed:?}; it should not depend on an occupied worker thread \
332 freeing up"
333 );
334
335 for occupier in occupiers {
336 occupier.await.unwrap();
337 }
338 }
339
340 #[tokio::test(flavor = "current_thread")]
343 async fn test_scope_current_thread_runtime() {
344 let results = tokio::task::spawn_blocking(|| {
345 scope_bounded(16, |scope| {
346 for i in 0..16 {
347 scope.spawn(move || i);
348 }
349 })
350 .collect::<Vec<_>>()
351 })
352 .await
353 .unwrap();
354 assert_eq!(results.len(), 16);
355 results.iter().enumerate().for_each(|(i, &result)| {
356 assert_eq!(result, i);
357 });
358 }
359
360 #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
363 async fn test_scope_runs_in_parallel() {
364 const JOBS: usize = 16;
365 const PER_JOB: Duration = Duration::from_millis(50);
366 let started = Instant::now();
367 let results = tokio::task::spawn_blocking(|| {
368 scope_bounded(JOBS, |scope| {
369 for i in 0..JOBS {
370 scope.spawn(move || {
371 thread::sleep(PER_JOB);
372 i
373 });
374 }
375 })
376 .collect::<Vec<_>>()
377 })
378 .await
379 .unwrap();
380 let elapsed = started.elapsed();
381 assert_eq!(results.len(), JOBS);
382 assert!(
385 elapsed < (JOBS as u32 * PER_JOB) / 2,
386 "scope_bounded took {elapsed:?}; expected parallel speedup across worker threads"
387 );
388 }
389
390 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
391 async fn test_scope() {
392 let results = scope_bounded(1000, |scope| {
393 for i in 0..1000 {
394 scope.spawn(move || i);
395 }
396 });
397 let results = results.collect::<Vec<_>>();
398 results.iter().enumerate().for_each(|(i, &result)| {
399 assert_eq!(result, i);
400 });
401 assert_eq!(results.len(), 1000);
402 }
403
404 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
405 async fn test_empty_scope() {
406 let results = scope_bounded(0, |scope| {
407 if false {
408 scope.spawn(|| 42);
409 }
410 });
411 assert_eq!(results.count(), 0);
412 }
413
414 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
415 async fn test_single_task() {
416 let results = scope_bounded(1, |scope| {
417 scope.spawn(|| 42);
418 })
419 .collect::<Vec<_>>();
420 assert_eq!(results, vec![42]);
421 }
422
423 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
424 async fn test_task_finish_before_scope() {
425 let results = scope_bounded(1, |scope| {
426 scope.spawn(|| 42);
427 thread::sleep(std::time::Duration::from_millis(100));
428 })
429 .collect::<Vec<_>>();
430 assert_eq!(results, vec![42]);
431 }
432
433 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
434 async fn test_task_finish_after_scope() {
435 let results = scope_bounded(1, |scope| {
436 scope.spawn(|| {
437 thread::sleep(std::time::Duration::from_millis(100));
438 42
439 });
440 })
441 .collect::<Vec<_>>();
442 assert_eq!(results, vec![42]);
443 }
444
445 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
446 async fn test_panic_in_scope_factory() {
447 let result = catch_unwind(AssertUnwindSafe(|| {
448 let _results = scope_bounded(1000, |scope| {
449 for i in 0..500 {
450 scope.spawn(move || i);
451 }
452 panic!("Intentional panic");
453 });
454 unreachable!();
455 }));
456 assert!(result.is_err());
457 assert_eq!(
458 result.unwrap_err().downcast_ref::<&str>(),
459 Some(&"Intentional panic")
460 );
461 }
462
463 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
464 async fn test_panic_in_scope_task() {
465 let result = catch_unwind(AssertUnwindSafe(|| {
466 let _results = scope_bounded(1000, |scope| {
467 for i in 0..1000 {
468 scope.spawn(move || {
469 if i == 500 {
470 panic!("Intentional panic");
471 } else if i == 501 {
472 panic!("Wrong intentional panic");
473 } else {
474 i
475 }
476 });
477 }
478 });
479 unreachable!();
480 }));
481 assert!(result.is_err());
482 assert_eq!(
483 result.unwrap_err().downcast_ref::<&str>(),
484 Some(&"Intentional panic")
485 );
486 }
487}