Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_text/tests/annealer_corpus/rayon_pool.rs

34.5 KiB, 1 run

created by r1870400018:11786, which is this file's identity for as long as the history lasts, whatever it is later renamed to

download · who wrote it · its history

1use crate::job::{JobFifo, JobRef, StackJob};
2use crate::latch::{AsCoreLatch, CoreLatch, Latch, LatchRef, LockLatch, OnceLatch, SpinLatch};
3use crate::sleep::Sleep;
4use crate::sync::Mutex;
5use crate::unwind;
6use crate::{
7 ErrorKind, ExitHandler, PanicHandler, StartHandler, ThreadPoolBuildError, ThreadPoolBuilder,
8 Yield,
9};
10use crossbeam_deque::{Injector, Steal, Stealer, Worker};
11use std::cell::Cell;
12use std::fmt;
13use std::hash::{DefaultHasher, Hasher};
14use std::io;
15use std::mem;
16use std::ptr;
17use std::sync::atomic::{AtomicUsize, Ordering};
18use std::sync::{Arc, Once};
19use std::thread;
20
21/// Thread builder used for customization via [`ThreadPoolBuilder::spawn_handler()`].
22pub struct ThreadBuilder {
23 name: Option<String>,
24 stack_size: Option<usize>,
25 worker: Worker<JobRef>,
26 stealer: Stealer<JobRef>,
27 registry: Arc<Registry>,
28 index: usize,
29}
30
31impl ThreadBuilder {
32 /// Gets the index of this thread in the pool, within `0..num_threads`.
33 pub fn index(&self) -> usize {
34 self.index
35 }
36
37 /// Gets the string that was specified by `ThreadPoolBuilder::name()`.
38 pub fn name(&self) -> Option<&str> {
39 self.name.as_deref()
40 }
41
42 /// Gets the value that was specified by `ThreadPoolBuilder::stack_size()`.
43 pub fn stack_size(&self) -> Option<usize> {
44 self.stack_size
45 }
46
47 /// Executes the main loop for this thread. This will not return until the
48 /// thread pool is dropped.
49 pub fn run(self) {
50 unsafe { main_loop(self) }
51 }
52}
53
54impl fmt::Debug for ThreadBuilder {
55 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
56 f.debug_struct("ThreadBuilder")
57 .field("pool", &self.registry.id())
58 .field("index", &self.index)
59 .field("name", &self.name)
60 .field("stack_size", &self.stack_size)
61 .finish()
62 }
63}
64
65/// Generalized trait for spawning a thread in the `Registry`.
66///
67/// This trait is crate-private, because we don't actually want to
68/// expose these details in the API.
69pub(crate) trait ThreadSpawn {
70 /// Spawn a thread with the `ThreadBuilder` parameters, and then
71 /// call `ThreadBuilder::run()`.
72 fn spawn(&mut self, thread: ThreadBuilder) -> io::Result<()>;
73}
74
75/// Spawns a thread in the "normal" way with `std::thread::Builder`.
76///
77/// This type is pub-in-private -- it has to be somewhat exposed to be
78/// usable as a type parameter of `ThreadPoolBuilder`, but we don't
79/// actually want to expose these details in the API.
80#[derive(Debug, Default)]
81#[expect(unnameable_types)]
82pub struct DefaultSpawn;
83
84impl ThreadSpawn for DefaultSpawn {
85 fn spawn(&mut self, thread: ThreadBuilder) -> io::Result<()> {
86 let mut b = thread::Builder::new();
87 if let Some(name) = thread.name() {
88 b = b.name(name.to_owned());
89 }
90 if let Some(stack_size) = thread.stack_size() {
91 b = b.stack_size(stack_size);
92 }
93 b.spawn(|| thread.run())?;
94 Ok(())
95 }
96}
97
98/// Spawns a thread with a user's custom callback.
99///
100/// This type is pub-in-private -- it has to be somewhat exposed to be
101/// usable as a type parameter of `ThreadPoolBuilder`, but we don't
102/// actually want to expose these details in the API.
103#[derive(Debug)]
104#[expect(unnameable_types)]
105pub struct CustomSpawn<F>(F);
106
107impl<F> CustomSpawn<F>
108where
109 F: FnMut(ThreadBuilder) -> io::Result<()>,
110{
111 pub(super) fn new(spawn: F) -> Self {
112 CustomSpawn(spawn)
113 }
114}
115
116impl<F> ThreadSpawn for CustomSpawn<F>
117where
118 F: FnMut(ThreadBuilder) -> io::Result<()>,
119{
120 #[inline]
121 fn spawn(&mut self, thread: ThreadBuilder) -> io::Result<()> {
122 (self.0)(thread)
123 }
124}
125
126pub(super) struct Registry {
127 thread_infos: Vec<ThreadInfo>,
128 sleep: Sleep,
129 injected_jobs: Injector<JobRef>,
130 broadcasts: Mutex<Vec<Worker<JobRef>>>,
131 panic_handler: Option<Box<PanicHandler>>,
132 start_handler: Option<Box<StartHandler>>,
133 exit_handler: Option<Box<ExitHandler>>,
134
135 // When this latch reaches 0, it means that all work on this
136 // registry must be complete. This is ensured in the following ways:
137 //
138 // - if this is the global registry, there is a ref-count that never
139 // gets released.
140 // - if this is a user-created thread pool, then so long as the thread pool
141 // exists, it holds a reference.
142 // - when we inject a "blocking job" into the registry with `ThreadPool::install()`,
143 // no adjustment is needed; the `ThreadPool` holds the reference, and since we won't
144 // return until the blocking job is complete, that ref will continue to be held.
145 // - when `join()` or `scope()` is invoked, similarly, no adjustments are needed.
146 // These are always owned by some other job (e.g., one injected by `ThreadPool::install()`)
147 // and that job will keep the pool alive.
148 terminate_count: AtomicUsize,
149}
150
151// ////////////////////////////////////////////////////////////////////////
152// Initialization
153
154static mut THE_REGISTRY: Option<Arc<Registry>> = None;
155static THE_REGISTRY_SET: Once = Once::new();
156
157/// Starts the worker threads (if that has not already happened). If
158/// initialization has not already occurred, use the default
159/// configuration.
160pub(super) fn global_registry() -> &'static Arc<Registry> {
161 set_global_registry(default_global_registry)
162 .or_else(|err| {
163 // SAFETY: we only create a shared reference to `THE_REGISTRY` after the `call_once`
164 // that initializes it, and there will be no more mutable accesses at all.
165 debug_assert!(THE_REGISTRY_SET.is_completed());
166 let the_registry = unsafe { &*ptr::addr_of!(THE_REGISTRY) };
167 the_registry.as_ref().ok_or(err)
168 })
169 .expect("The global thread pool has not been initialized.")
170}
171
172/// Starts the worker threads (if that has not already happened) with
173/// the given builder.
174pub(super) fn init_global_registry<S>(
175 builder: ThreadPoolBuilder<S>,
176) -> Result<&'static Arc<Registry>, ThreadPoolBuildError>
177where
178 S: ThreadSpawn,
179{
180 set_global_registry(|| Registry::new(builder))
181}
182
183/// Starts the worker threads (if that has not already happened)
184/// by creating a registry with the given callback.
185fn set_global_registry<F>(registry: F) -> Result<&'static Arc<Registry>, ThreadPoolBuildError>
186where
187 F: FnOnce() -> Result<Arc<Registry>, ThreadPoolBuildError>,
188{
189 let mut result = Err(ThreadPoolBuildError::new(
190 ErrorKind::GlobalPoolAlreadyInitialized,
191 ));
192
193 THE_REGISTRY_SET.call_once(|| {
194 result = registry().map(|registry: Arc<Registry>| {
195 // SAFETY: this is the only mutable access to `THE_REGISTRY`, thanks to `Once`, and
196 // `global_registry()` only takes a shared reference **after** this `call_once`.
197 unsafe {
198 ptr::addr_of_mut!(THE_REGISTRY).write(Some(registry));
199 (*ptr::addr_of!(THE_REGISTRY)).as_ref().unwrap_unchecked()
200 }
201 })
202 });
203
204 result
205}
206
207fn default_global_registry() -> Result<Arc<Registry>, ThreadPoolBuildError> {
208 let result = Registry::new(ThreadPoolBuilder::new());
209
210 // If we're running in an environment that doesn't support threads at all, we can fall back to
211 // using the current thread alone. This is crude, and probably won't work for non-blocking
212 // calls like `spawn` or `broadcast_spawn`, but a lot of stuff does work fine.
213 //
214 // Notably, this allows current WebAssembly targets to work even though their threading support
215 // is stubbed out, and we won't have to change anything if they do add real threading.
216 let unsupported = matches!(&result, Err(e) if e.is_unsupported());
217 if unsupported && WorkerThread::current().is_null() {
218 let builder = ThreadPoolBuilder::new().num_threads(1).use_current_thread();
219 let fallback_result = Registry::new(builder);
220 if fallback_result.is_ok() {
221 return fallback_result;
222 }
223 }
224
225 result
226}
227
228struct Terminator<'a>(&'a Arc<Registry>);
229
230impl<'a> Drop for Terminator<'a> {
231 fn drop(&mut self) {
232 self.0.terminate()
233 }
234}
235
236impl Registry {
237 pub(super) fn new<S>(
238 mut builder: ThreadPoolBuilder<S>,
239 ) -> Result<Arc<Self>, ThreadPoolBuildError>
240 where
241 S: ThreadSpawn,
242 {
243 // Soft-limit the number of threads that we can actually support.
244 let n_threads = Ord::min(builder.get_num_threads(), crate::max_num_threads());
245
246 let breadth_first = builder.get_breadth_first();
247
248 let (workers, stealers): (Vec<_>, Vec<_>) = (0..n_threads)
249 .map(|_| {
250 let worker = if breadth_first {
251 Worker::new_fifo()
252 } else {
253 Worker::new_lifo()
254 };
255
256 let stealer = worker.stealer();
257 (worker, stealer)
258 })
259 .unzip();
260
261 let (broadcasts, broadcast_stealers): (Vec<_>, Vec<_>) = (0..n_threads)
262 .map(|_| {
263 let worker = Worker::new_fifo();
264 let stealer = worker.stealer();
265 (worker, stealer)
266 })
267 .unzip();
268
269 let registry = Arc::new(Registry {
270 thread_infos: stealers.into_iter().map(ThreadInfo::new).collect(),
271 sleep: Sleep::new(n_threads),
272 injected_jobs: Injector::new(),
273 broadcasts: Mutex::new(broadcasts),
274 terminate_count: AtomicUsize::new(1),
275 panic_handler: builder.take_panic_handler(),
276 start_handler: builder.take_start_handler(),
277 exit_handler: builder.take_exit_handler(),
278 });
279
280 // If we return early or panic, make sure to terminate existing threads.
281 let t1000 = Terminator(&registry);
282
283 for (index, (worker, stealer)) in workers.into_iter().zip(broadcast_stealers).enumerate() {
284 let thread = ThreadBuilder {
285 name: builder.get_thread_name(index),
286 stack_size: builder.get_stack_size(),
287 registry: Arc::clone(&registry),
288 worker,
289 stealer,
290 index,
291 };
292
293 if index == 0 && builder.use_current_thread {
294 if !WorkerThread::current().is_null() {
295 return Err(ThreadPoolBuildError::new(
296 ErrorKind::CurrentThreadAlreadyInPool,
297 ));
298 }
299 // Rather than starting a new thread, we're just taking over the current thread
300 // *without* running the main loop, so we can still return from here.
301 // The WorkerThread is leaked, but we never shutdown the global pool anyway.
302 let worker_thread = Box::into_raw(Box::new(WorkerThread::from(thread)));
303
304 unsafe {
305 WorkerThread::set_current(worker_thread);
306 Latch::set(&registry.thread_infos[index].primed);
307 }
308 continue;
309 }
310
311 if let Err(e) = builder.get_spawn_handler().spawn(thread) {
312 return Err(ThreadPoolBuildError::new(ErrorKind::IOError(e)));
313 }
314 }
315
316 // Returning normally now, without termination.
317 mem::forget(t1000);
318
319 Ok(registry)
320 }
321
322 pub(super) fn current() -> Arc<Registry> {
323 unsafe {
324 let worker_thread = WorkerThread::current();
325 let registry = if worker_thread.is_null() {
326 global_registry()
327 } else {
328 &(*worker_thread).registry
329 };
330 Arc::clone(registry)
331 }
332 }
333
334 /// Returns the number of threads in the current registry. This
335 /// is better than `Registry::current().num_threads()` because it
336 /// avoids incrementing the `Arc`.
337 pub(super) fn current_num_threads() -> usize {
338 unsafe {
339 let worker_thread = WorkerThread::current();
340 if worker_thread.is_null() {
341 global_registry().num_threads()
342 } else {
343 (*worker_thread).registry.num_threads()
344 }
345 }
346 }
347
348 /// Returns the current `WorkerThread` if it's part of this `Registry`.
349 pub(super) fn current_thread(&self) -> Option<&WorkerThread> {
350 unsafe {
351 let worker = WorkerThread::current().as_ref()?;
352 if worker.registry().id() == self.id() {
353 Some(worker)
354 } else {
355 None
356 }
357 }
358 }
359
360 /// Returns an opaque identifier for this registry.
361 pub(super) fn id(&self) -> RegistryId {
362 // We can rely on `self` not to change since we only ever create
363 // registries that are boxed up in an `Arc` (see `new()` above).
364 RegistryId {
365 addr: self as *const Self as usize,
366 }
367 }
368
369 pub(super) fn num_threads(&self) -> usize {
370 self.thread_infos.len()
371 }
372
373 pub(super) fn catch_unwind(&self, f: impl FnOnce()) {
374 if let Err(err) = unwind::halt_unwinding(f) {
375 // If there is no handler, or if that handler itself panics, then we abort.
376 let abort_guard = unwind::AbortIfPanic;
377 if let Some(ref handler) = self.panic_handler {
378 handler(err);
379 mem::forget(abort_guard);
380 }
381 }
382 }
383
384 /// Waits for the worker threads to get up and running. This is
385 /// meant to be used for benchmarking purposes, primarily, so that
386 /// you can get more consistent numbers by having everything
387 /// "ready to go".
388 pub(super) fn wait_until_primed(&self) {
389 for info in &self.thread_infos {
390 info.primed.wait();
391 }
392 }
393
394 /// Waits for the worker threads to stop. This is used for testing
395 /// -- so we can check that termination actually works.
396 #[cfg(test)]
397 pub(super) fn wait_until_stopped(&self) {
398 for info in &self.thread_infos {
399 info.stopped.wait();
400 }
401 }
402
403 // ////////////////////////////////////////////////////////////////////////
404 // MAIN LOOP
405 //
406 // So long as all of the worker threads are hanging out in their
407 // top-level loop, there is no work to be done.
408
409 /// Push a job into the given `registry`. If we are running on a
410 /// worker thread for the registry, this will push onto the
411 /// deque. Else, it will inject from the outside (which is slower).
412 pub(super) fn inject_or_push(&self, job_ref: JobRef) {
413 let worker_thread = WorkerThread::current();
414 unsafe {
415 if !worker_thread.is_null() && (*worker_thread).registry().id() == self.id() {
416 (*worker_thread).push(job_ref);
417 } else {
418 self.inject(job_ref);
419 }
420 }
421 }
422
423 /// Push a job into the "external jobs" queue; it will be taken by
424 /// whatever worker has nothing to do. Use this if you know that
425 /// you are not on a worker of this registry.
426 pub(super) fn inject(&self, injected_job: JobRef) {
427 // It should not be possible for `state.terminate` to be true
428 // here. It is only set to true when the user creates (and
429 // drops) a `ThreadPool`; and, in that case, they cannot be
430 // calling `inject()` later, since they dropped their
431 // `ThreadPool`.
432 debug_assert_ne!(
433 self.terminate_count.load(Ordering::Acquire),
434 0,
435 "inject() sees state.terminate as true"
436 );
437
438 let queue_was_empty = self.injected_jobs.is_empty();
439
440 self.injected_jobs.push(injected_job);
441 self.sleep.new_injected_jobs(1, queue_was_empty);
442 }
443
444 fn has_injected_job(&self) -> bool {
445 !self.injected_jobs.is_empty()
446 }
447
448 fn pop_injected_job(&self) -> Option<JobRef> {
449 loop {
450 match self.injected_jobs.steal() {
451 Steal::Success(job) => return Some(job),
452 Steal::Empty => return None,
453 Steal::Retry => {}
454 }
455 }
456 }
457
458 /// Push a job into each thread's own "external jobs" queue; it will be
459 /// executed only on that thread, when it has nothing else to do locally,
460 /// before it tries to steal other work.
461 ///
462 /// **Panics** if not given exactly as many jobs as there are threads.
463 pub(super) fn inject_broadcast(&self, injected_jobs: impl ExactSizeIterator<Item = JobRef>) {
464 assert_eq!(self.num_threads(), injected_jobs.len());
465 {
466 let broadcasts = self.broadcasts.lock().unwrap();
467
468 // It should not be possible for `state.terminate` to be true
469 // here. It is only set to true when the user creates (and
470 // drops) a `ThreadPool`; and, in that case, they cannot be
471 // calling `inject_broadcast()` later, since they dropped their
472 // `ThreadPool`.
473 debug_assert_ne!(
474 self.terminate_count.load(Ordering::Acquire),
475 0,
476 "inject_broadcast() sees state.terminate as true"
477 );
478
479 assert_eq!(broadcasts.len(), injected_jobs.len());
480 for (worker, job_ref) in broadcasts.iter().zip(injected_jobs) {
481 worker.push(job_ref);
482 }
483 }
484 for i in 0..self.num_threads() {
485 self.sleep.notify_worker_latch_is_set(i);
486 }
487 }
488
489 /// If already in a worker-thread of this registry, just execute `op`.
490 /// Otherwise, inject `op` in this thread pool. Either way, block until `op`
491 /// completes and return its return value. If `op` panics, that panic will
492 /// be propagated as well. The second argument indicates `true` if injection
493 /// was performed, `false` if executed directly.
494 pub(super) fn in_worker<OP, R>(&self, op: OP) -> R
495 where
496 OP: FnOnce(&WorkerThread, bool) -> R + Send,
497 R: Send,
498 {
499 unsafe {
500 let worker_thread = WorkerThread::current();
501 if worker_thread.is_null() {
502 self.in_worker_cold(op)
503 } else if (*worker_thread).registry().id() != self.id() {
504 self.in_worker_cross(&*worker_thread, op)
505 } else {
506 // Perfectly valid to give them a `&T`: this is the
507 // current thread, so we know the data structure won't be
508 // invalidated until we return.
509 op(&*worker_thread, false)
510 }
511 }
512 }
513
514 #[cold]
515 unsafe fn in_worker_cold<OP, R>(&self, op: OP) -> R
516 where
517 OP: FnOnce(&WorkerThread, bool) -> R + Send,
518 R: Send,
519 {
520 thread_local!(static LOCK_LATCH: LockLatch = const { LockLatch::new() });
521
522 LOCK_LATCH.with(|l| unsafe {
523 // This thread isn't a member of *any* thread pool, so just block.
524 debug_assert!(WorkerThread::current().is_null());
525 let job = StackJob::new(
526 |injected| {
527 let worker_thread = WorkerThread::current();
528 assert!(injected && !worker_thread.is_null());
529 op(&*worker_thread, true)
530 },
531 LatchRef::new(l),
532 );
533 self.inject(job.as_job_ref());
534 job.latch.wait_and_reset(); // Make sure we can use the same latch again next time.
535
536 job.into_result()
537 })
538 }
539
540 #[cold]
541 unsafe fn in_worker_cross<OP, R>(&self, current_thread: &WorkerThread, op: OP) -> R
542 where
543 OP: FnOnce(&WorkerThread, bool) -> R + Send,
544 R: Send,
545 {
546 unsafe {
547 // This thread is a member of a different pool, so let it process
548 // other work while waiting for this `op` to complete.
549 debug_assert!(current_thread.registry().id() != self.id());
550 let latch = SpinLatch::cross(current_thread);
551 let job = StackJob::new(
552 |injected| {
553 let worker_thread = WorkerThread::current();
554 assert!(injected && !worker_thread.is_null());
555 op(&*worker_thread, true)
556 },
557 latch,
558 );
559 self.inject(job.as_job_ref());
560 current_thread.wait_until(&job.latch);
561 job.into_result()
562 }
563 }
564
565 /// Increments the terminate counter. This increment should be
566 /// balanced by a call to `terminate`, which will decrement. This
567 /// is used when spawning asynchronous work, which needs to
568 /// prevent the registry from terminating so long as it is active.
569 ///
570 /// Note that blocking functions such as `join` and `scope` do not
571 /// need to concern themselves with this fn; their context is
572 /// responsible for ensuring the current thread pool will not
573 /// terminate until they return.
574 ///
575 /// The global thread pool always has an outstanding reference
576 /// (the initial one). Custom thread pools have one outstanding
577 /// reference that is dropped when the `ThreadPool` is dropped:
578 /// since installing the thread pool blocks until any joins/scopes
579 /// complete, this ensures that joins/scopes are covered.
580 ///
581 /// The exception is `::spawn()`, which can create a job outside
582 /// of any blocking scope. In that case, the job itself holds a
583 /// terminate count and is responsible for invoking `terminate()`
584 /// when finished.
585 pub(super) fn increment_terminate_count(&self) {
586 let previous = self.terminate_count.fetch_add(1, Ordering::AcqRel);
587 debug_assert!(previous != 0, "registry ref count incremented from zero");
588 assert!(previous != usize::MAX, "overflow in registry ref count");
589 }
590
591 /// Signals that the thread pool which owns this registry has been
592 /// dropped. The worker threads will gradually terminate, once any
593 /// extant work is completed.
594 pub(super) fn terminate(&self) {
595 if self.terminate_count.fetch_sub(1, Ordering::AcqRel) == 1 {
596 for (i, thread_info) in self.thread_infos.iter().enumerate() {
597 unsafe { OnceLatch::set_and_tickle_one(&thread_info.terminate, self, i) };
598 }
599 }
600 }
601
602 /// Notify the worker that the latch they are sleeping on has been "set".
603 pub(super) fn notify_worker_latch_is_set(&self, target_worker_index: usize) {
604 self.sleep.notify_worker_latch_is_set(target_worker_index);
605 }
606}
607
608#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
609pub(super) struct RegistryId {
610 addr: usize,
611}
612
613struct ThreadInfo {
614 /// Latch set once thread has started and we are entering into the
615 /// main loop. Used to wait for worker threads to become primed,
616 /// primarily of interest for benchmarking.
617 primed: LockLatch,
618
619 /// Latch is set once worker thread has completed. Used to wait
620 /// until workers have stopped; only used for tests.
621 stopped: LockLatch,
622
623 /// The latch used to signal that terminated has been requested.
624 /// This latch is *set* by the `terminate` method on the
625 /// `Registry`, once the registry's main "terminate" counter
626 /// reaches zero.
627 terminate: OnceLatch,
628
629 /// the "stealer" half of the worker's deque
630 stealer: Stealer<JobRef>,
631}
632
633impl ThreadInfo {
634 fn new(stealer: Stealer<JobRef>) -> ThreadInfo {
635 ThreadInfo {
636 primed: LockLatch::new(),
637 stopped: LockLatch::new(),
638 terminate: OnceLatch::new(),
639 stealer,
640 }
641 }
642}
643
644// ////////////////////////////////////////////////////////////////////////
645// WorkerThread identifiers
646
647pub(super) struct WorkerThread {
648 /// the "worker" half of our local deque
649 worker: Worker<JobRef>,
650
651 /// the "stealer" half of the worker's broadcast deque
652 stealer: Stealer<JobRef>,
653
654 /// local queue used for `spawn_fifo` indirection
655 fifo: JobFifo,
656
657 index: usize,
658
659 /// A weak random number generator.
660 rng: XorShift64Star,
661
662 registry: Arc<Registry>,
663}
664
665// This is a bit sketchy, but basically: the WorkerThread is
666// allocated on the stack of the worker on entry and stored into this
667// thread-local variable. So it will remain valid at least until the
668// worker is fully unwound. Using an unsafe pointer avoids the need
669// for a RefCell<T> etc.
670thread_local! {
671 static WORKER_THREAD_STATE: Cell<*const WorkerThread> = const { Cell::new(ptr::null()) };
672}
673
674impl From<ThreadBuilder> for WorkerThread {
675 fn from(thread: ThreadBuilder) -> Self {
676 Self {
677 worker: thread.worker,
678 stealer: thread.stealer,
679 fifo: JobFifo::new(),
680 index: thread.index,
681 rng: XorShift64Star::new(),
682 registry: thread.registry,
683 }
684 }
685}
686
687impl Drop for WorkerThread {
688 fn drop(&mut self) {
689 // Undo `set_current`
690 WORKER_THREAD_STATE.with(|t| {
691 assert!(t.get().eq(&(self as *const _)));
692 t.set(ptr::null());
693 });
694 }
695}
696
697impl WorkerThread {
698 /// Gets the `WorkerThread` index for the current thread; returns
699 /// NULL if this is not a worker thread. This pointer is valid
700 /// anywhere on the current thread.
701 #[inline]
702 pub(super) fn current() -> *const WorkerThread {
703 WORKER_THREAD_STATE.get()
704 }
705
706 /// Sets `self` as the worker-thread index for the current thread.
707 /// This is done during worker-thread startup.
708 unsafe fn set_current(thread: *const WorkerThread) {
709 WORKER_THREAD_STATE.with(|t| {
710 assert!(t.get().is_null());
711 t.set(thread);
712 });
713 }
714
715 /// Returns the registry that owns this worker thread.
716 #[inline]
717 pub(super) fn registry(&self) -> &Arc<Registry> {
718 &self.registry
719 }
720
721 /// Our index amongst the worker threads (ranges from `0..self.num_threads()`).
722 #[inline]
723 pub(super) fn index(&self) -> usize {
724 self.index
725 }
726
727 #[inline]
728 pub(super) unsafe fn push(&self, job: JobRef) {
729 let queue_was_empty = self.worker.is_empty();
730 self.worker.push(job);
731 self.registry.sleep.new_internal_jobs(1, queue_was_empty);
732 }
733
734 #[inline]
735 pub(super) unsafe fn push_fifo(&self, job: JobRef) {
736 unsafe { self.push(self.fifo.push(job)) };
737 }
738
739 #[inline]
740 pub(super) fn local_deque_is_empty(&self) -> bool {
741 self.worker.is_empty()
742 }
743
744 /// Attempts to obtain a "local" job -- typically this means
745 /// popping from the top of the stack, though if we are configured
746 /// for breadth-first execution, it would mean dequeuing from the
747 /// bottom.
748 #[inline]
749 pub(super) fn take_local_job(&self) -> Option<JobRef> {
750 let popped_job = self.worker.pop();
751
752 if popped_job.is_some() {
753 return popped_job;
754 }
755
756 loop {
757 match self.stealer.steal() {
758 Steal::Success(job) => return Some(job),
759 Steal::Empty => return None,
760 Steal::Retry => {}
761 }
762 }
763 }
764
765 fn has_injected_job(&self) -> bool {
766 !self.stealer.is_empty() || self.registry.has_injected_job()
767 }
768
769 /// Wait until the latch is set. Try to keep busy by popping and
770 /// stealing tasks as necessary.
771 #[inline]
772 pub(super) unsafe fn wait_until<L: AsCoreLatch + ?Sized>(&self, latch: &L) {
773 let latch = latch.as_core_latch();
774 if !latch.probe() {
775 unsafe { self.wait_until_cold(latch) };
776 }
777 }
778
779 #[cold]
780 unsafe fn wait_until_cold(&self, latch: &CoreLatch) {
781 // the code below should swallow all panics and hence never
782 // unwind; but if something does wrong, we want to abort,
783 // because otherwise other code in rayon may assume that the
784 // latch has been signaled, and that can lead to random memory
785 // accesses, which would be *very bad*
786 let abort_guard = unwind::AbortIfPanic;
787
788 'outer: while !latch.probe() {
789 // Check for local work *before* we start marking ourself idle,
790 // especially to avoid modifying shared sleep state.
791 if let Some(job) = self.take_local_job() {
792 unsafe { self.execute(job) };
793 continue;
794 }
795
796 let mut idle_state = self.registry.sleep.start_looking(self.index);
797 while !latch.probe() {
798 if let Some(job) = self.find_work() {
799 self.registry.sleep.work_found();
800 unsafe { self.execute(job) };
801 // The job might have injected local work, so go back to the outer loop.
802 continue 'outer;
803 } else {
804 self.registry
805 .sleep
806 .no_work_found(&mut idle_state, latch, || self.has_injected_job())
807 }
808 }
809
810 // If we were sleepy, we are not anymore. We "found work" --
811 // whatever the surrounding thread was doing before it had to wait.
812 self.registry.sleep.work_found();
813 break;
814 }
815
816 mem::forget(abort_guard); // successful execution, do not abort
817 }
818
819 unsafe fn wait_until_out_of_work(&self) {
820 unsafe {
821 debug_assert_eq!(self as *const _, WorkerThread::current());
822 let registry = &*self.registry;
823 let index = self.index;
824
825 self.wait_until(&registry.thread_infos[index].terminate);
826
827 // Should not be any work left in our queue.
828 debug_assert!(self.take_local_job().is_none());
829
830 // Let registry know we are done
831 Latch::set(&registry.thread_infos[index].stopped);
832 }
833 }
834
835 fn find_work(&self) -> Option<JobRef> {
836 // Try to find some work to do. We give preference first
837 // to things in our local deque, then in other workers
838 // deques, and finally to injected jobs from the
839 // outside. The idea is to finish what we started before
840 // we take on something new.
841 self.take_local_job()
842 .or_else(|| self.steal())
843 .or_else(|| self.registry.pop_injected_job())
844 }
845
846 pub(super) fn yield_now(&self) -> Yield {
847 match self.find_work() {
848 Some(job) => unsafe {
849 self.execute(job);
850 Yield::Executed
851 },
852 None => Yield::Idle,
853 }
854 }
855
856 pub(super) fn yield_local(&self) -> Yield {
857 match self.take_local_job() {
858 Some(job) => unsafe {
859 self.execute(job);
860 Yield::Executed
861 },
862 None => Yield::Idle,
863 }
864 }
865
866 #[inline]
867 pub(super) unsafe fn execute(&self, job: JobRef) {
868 unsafe { job.execute() };
869 }
870
871 /// Try to steal a single job and return it.
872 ///
873 /// This should only be done as a last resort, when there is no
874 /// local work to do.
875 fn steal(&self) -> Option<JobRef> {
876 // we only steal when we don't have any work to do locally
877 debug_assert!(self.local_deque_is_empty());
878
879 // otherwise, try to steal
880 let thread_infos = &self.registry.thread_infos.as_slice();
881 let num_threads = thread_infos.len();
882 if num_threads <= 1 {
883 return None;
884 }
885
886 loop {
887 let mut retry = false;
888 let start = self.rng.next_usize(num_threads);
889 let job = (start..num_threads)
890 .chain(0..start)
891 .filter(move |&i| i != self.index)
892 .find_map(|victim_index| {
893 let victim = &thread_infos[victim_index];
894 match victim.stealer.steal() {
895 Steal::Success(job) => Some(job),
896 Steal::Empty => None,
897 Steal::Retry => {
898 retry = true;
899 None
900 }
901 }
902 });
903 if job.is_some() || !retry {
904 return job;
905 }
906 }
907 }
908}
909
910// ////////////////////////////////////////////////////////////////////////
911
912unsafe fn main_loop(thread: ThreadBuilder) {
913 unsafe {
914 let worker_thread = &WorkerThread::from(thread);
915 WorkerThread::set_current(worker_thread);
916 let registry = &*worker_thread.registry;
917 let index = worker_thread.index;
918
919 // let registry know we are ready to do work
920 Latch::set(&registry.thread_infos[index].primed);
921
922 // Worker threads should not panic. If they do, just abort, as the
923 // internal state of the thread pool is corrupted. Note that if
924 // **user code** panics, we should catch that and redirect.
925 let abort_guard = unwind::AbortIfPanic;
926
927 // Inform a user callback that we started a thread.
928 if let Some(ref handler) = registry.start_handler {
929 registry.catch_unwind(|| handler(index));
930 }
931
932 worker_thread.wait_until_out_of_work();
933
934 // Normal termination, do not abort.
935 mem::forget(abort_guard);
936
937 // Inform a user callback that we exited a thread.
938 if let Some(ref handler) = registry.exit_handler {
939 registry.catch_unwind(|| handler(index));
940 // We're already exiting the thread, there's nothing else to do.
941 }
942 }
943}
944
945/// If already in a worker-thread, just execute `op`. Otherwise,
946/// execute `op` in the default thread pool. Either way, block until
947/// `op` completes and return its return value. If `op` panics, that
948/// panic will be propagated as well. The second argument indicates
949/// `true` if injection was performed, `false` if executed directly.
950pub(super) fn in_worker<OP, R>(op: OP) -> R
951where
952 OP: FnOnce(&WorkerThread, bool) -> R + Send,
953 R: Send,
954{
955 unsafe {
956 let owner_thread = WorkerThread::current();
957 if !owner_thread.is_null() {
958 // Perfectly valid to give them a `&T`: this is the
959 // current thread, so we know the data structure won't be
960 // invalidated until we return.
961 op(&*owner_thread, false)
962 } else {
963 global_registry().in_worker(op)
964 }
965 }
966}
967
968/// [xorshift*] is a fast pseudorandom number generator which will
969/// even tolerate weak seeding, as long as it's not zero.
970///
971/// [xorshift*]: https://en.wikipedia.org/wiki/Xorshift#xorshift*
972struct XorShift64Star {
973 state: Cell<u64>,
974}
975
976impl XorShift64Star {
977 fn new() -> Self {
978 // Any non-zero seed will do -- this uses the hash of a global counter.
979 let mut seed = 0;
980 while seed == 0 {
981 let mut hasher = DefaultHasher::new();
982 static COUNTER: AtomicUsize = AtomicUsize::new(0);
983 hasher.write_usize(COUNTER.fetch_add(1, Ordering::Relaxed));
984 seed = hasher.finish();
985 }
986
987 XorShift64Star {
988 state: Cell::new(seed),
989 }
990 }
991
992 fn next(&self) -> u64 {
993 let mut x = self.state.get();
994 debug_assert_ne!(x, 0);
995 x ^= x >> 12;
996 x ^= x << 25;
997 x ^= x >> 27;
998 self.state.set(x);
999 x.wrapping_mul(0x2545_f491_4f6c_dd1d)
1000 }
1001
1002 /// Return a value from `0..n`.
1003 fn next_usize(&self, n: usize) -> usize {
1004 (self.next() % n as u64) as usize
1005 }
1006}