Skip to main content

tenferro_cpu/
context.rs

1use std::env;
2use std::num::NonZeroUsize;
3use std::sync::Arc;
4
5use rayon::prelude::*;
6use thiserror::Error as ThisError;
7
8use crate::affinity::{CpuAffinityError, SystemThreadAffinity, ThreadAffinity};
9use crate::arbiter::{
10    current_execution_owner, register_worker_execution_scope, worker_execution_scope_matches,
11    ExecutionScopeState,
12};
13use crate::domain_executor::{
14    CpuDomainExecutor, CpuDomainExecutorCapabilities, CpuDomainExecutorError, CpuExecutorAffinity,
15    CpuExecutorReentrancy, CpuExecutorShutdown, CpuInnerParallelism, ScopedCpuJob, ScopedCpuJobs,
16};
17use crate::{CpuSet, Error, ErrorKind, Result, ValidationKind};
18
19/// Stack size reserved for every Tenferro CPU worker thread.
20///
21/// Rust's `std::thread` default is 2 MiB, but provider code can recurse with
22/// large private frames: a `NUM_THREADS=64` OpenBLAS build keeps a
23/// `job_t job[64]` (about 541 KiB per level, measured) on the stack of every
24/// recursive `dgetrf_parallel` frame, so 2 MiB allowed only three levels and
25/// aborted with a stack overflow at n>=256. Sixteen MiB leaves room for about
26/// thirty such frames, which is the order of the 8 MiB main-thread default that
27/// provider calls receive outside a pool.
28///
29/// Override it per context with [`CpuContext::with_threads_and_worker_stack`],
30/// or for the environment-configured path with the
31/// `TENFERRO_CPU_WORKER_STACK_BYTES` environment variable.
32pub const DEFAULT_WORKER_STACK_BYTES: usize = 16 << 20;
33
34/// Smallest accepted worker stack size.
35///
36/// A smaller stack cannot run a nontrivial provider call, so rejecting it while
37/// configuring the pool replaces an eventual stack-overflow abort with a typed
38/// configuration error.
39const MIN_WORKER_STACK_BYTES: usize = 64 << 10;
40
41/// Resolve the worker stack size, honoring `TENFERRO_CPU_WORKER_STACK_BYTES`.
42///
43/// A malformed or unreadable variable is a configuration error rather than a
44/// silent fallback, so a deployment that configures the wrong value finds out at
45/// context construction instead of through an eventual stack overflow.
46fn worker_stack_bytes_from_env() -> Result<usize> {
47    match env::var("TENFERRO_CPU_WORKER_STACK_BYTES") {
48        Ok(value) => value.parse::<usize>().map_err(|err| {
49            Error::extension(
50                "CpuContext::with_threads",
51                "cpu",
52                ErrorKind::Validation(ValidationKind::InvalidArgument),
53                err,
54            )
55        }),
56        Err(env::VarError::NotPresent) => Ok(DEFAULT_WORKER_STACK_BYTES),
57        Err(err) => Err(Error::extension(
58            "CpuContext::with_threads",
59            "cpu",
60            ErrorKind::Validation(ValidationKind::InvalidArgument),
61            err,
62        )),
63    }
64}
65
66/// Failure to construct a CPU context with pinned Rayon workers.
67///
68/// # Examples
69///
70/// ```
71/// use tenferro_cpu::CpuContextError;
72///
73/// let error = CpuContextError::InvalidThreadCount;
74/// assert!(error.to_string().contains("thread count"));
75/// ```
76#[derive(Debug, ThisError)]
77pub enum CpuContextError {
78    /// A context must contain at least one worker.
79    #[error("thread count must be at least 1")]
80    InvalidThreadCount,
81    /// A pinned engine cannot create more workers than assigned CPUs.
82    #[error("requested {workers} workers for only {cpus} assigned CPUs")]
83    TooManyWorkers {
84        /// Requested Rayon worker count.
85        workers: usize,
86        /// Number of logical CPUs in the execution domain.
87        cpus: usize,
88    },
89    /// The environment worker stack configuration was rejected.
90    #[error("invalid worker stack configuration")]
91    InvalidWorkerStack {
92        /// Underlying configuration failure.
93        #[source]
94        source: Error,
95    },
96    /// Rayon could not construct the custom thread pool.
97    #[error("failed to build pinned CPU thread pool: {source}")]
98    PoolBuild {
99        /// Rayon or OS thread-spawn error.
100        #[source]
101        source: rayon::ThreadPoolBuildError,
102    },
103    /// A worker could not be confined to the domain CPU set.
104    #[error("failed to confine worker {worker} to the domain CPU set: {source}")]
105    WorkerAffinity {
106        /// Stable Rayon worker index.
107        worker: usize,
108        /// OS or verification failure.
109        #[source]
110        source: CpuAffinityError,
111    },
112    /// A worker terminated before reporting startup affinity.
113    #[error("worker startup channel closed before all workers reported: {source}")]
114    WorkerStartupClosed {
115        /// Channel receive failure from the worker startup handshake.
116        #[source]
117        source: std::sync::mpsc::RecvError,
118    },
119}
120
121/// Reusable CPU execution context carrying CPU parallelism policy.
122///
123/// `CpuContext` stores the requested thread count as a kernel-level
124/// parallelism hint and owns the Rayon pool used by multi-threaded CPU work.
125///
126/// # Examples
127///
128/// ```
129/// use tenferro_cpu::CpuContext;
130///
131/// let ctx = CpuContext::with_threads(1).unwrap();
132/// let value = ctx.install(|| 1 + 1);
133/// assert_eq!(value, 2);
134/// assert_eq!(ctx.num_threads(), 1);
135/// ```
136#[derive(Clone, Debug)]
137pub struct CpuContext {
138    num_threads: usize,
139    worker_stack_bytes: usize,
140    pool: Option<Arc<rayon::ThreadPool>>,
141    pinned_cpus: Option<CpuSet>,
142    execution_scope: Arc<ExecutionScopeState>,
143    #[cfg(test)]
144    executor_install_calls: Arc<std::sync::atomic::AtomicUsize>,
145}
146
147impl CpuContext {
148    /// Create a CPU context from `RAYON_NUM_THREADS`, or fall back to a
149    /// single-threaded context with a stderr warning when validation fails.
150    ///
151    /// # Examples
152    ///
153    /// ```
154    /// use tenferro_cpu::CpuContext;
155    ///
156    /// let ctx = CpuContext::from_env();
157    /// assert!(ctx.num_threads() >= 1);
158    /// ```
159    pub fn from_env() -> Self {
160        Self::try_from_env().unwrap_or_else(|err| {
161            eprintln!(
162                "tenferro_cpu: falling back to single-threaded CPU context after configuration error: {err}"
163            );
164            Self::single_threaded()
165        })
166    }
167
168    /// Try to create a CPU context from `RAYON_NUM_THREADS`.
169    ///
170    /// # Examples
171    ///
172    /// ```
173    /// use tenferro_cpu::CpuContext;
174    ///
175    /// let ctx = CpuContext::try_from_env()
176    ///     .unwrap_or_else(|_| CpuContext::with_threads(1).unwrap());
177    /// assert!(ctx.num_threads() >= 1);
178    /// ```
179    ///
180    /// # Errors
181    ///
182    /// Returns [`CpuContextError`] when `RAYON_NUM_THREADS` is malformed or
183    /// requests an invalid worker count.
184    pub fn try_from_env() -> Result<Self> {
185        match env::var("RAYON_NUM_THREADS") {
186            Ok(value) => {
187                let num_threads = value.parse::<usize>().map_err(|err| {
188                    Error::extension(
189                        "CpuContext::try_from_env",
190                        "cpu",
191                        ErrorKind::Validation(ValidationKind::InvalidArgument),
192                        err,
193                    )
194                })?;
195                Self::with_threads_and_worker_stack(num_threads, worker_stack_bytes_from_env()?)
196                    .map_err(|err| match err {
197                        Error::Validation { source, .. } => {
198                            Error::validation("CpuContext::try_from_env", source)
199                        }
200                        err => err,
201                    })
202            }
203            Err(env::VarError::NotPresent) => {
204                Self::with_threads(super::affinity::available_parallelism())
205            }
206            Err(err) => Err(Error::extension(
207                "CpuContext::try_from_env",
208                "cpu",
209                ErrorKind::Validation(ValidationKind::InvalidArgument),
210                err,
211            )),
212        }
213    }
214
215    /// Create a CPU context with a fixed parallelism hint.
216    ///
217    /// The worker stack size comes from `TENFERRO_CPU_WORKER_STACK_BYTES` when
218    /// that variable is present, and from [`DEFAULT_WORKER_STACK_BYTES`]
219    /// otherwise. Use [`CpuContext::with_threads_and_worker_stack`] to choose it
220    /// programmatically.
221    ///
222    /// # Examples
223    ///
224    /// ```
225    /// use tenferro_cpu::CpuContext;
226    ///
227    /// let ctx = CpuContext::with_threads(2).unwrap();
228    /// assert_eq!(ctx.num_threads(), 2);
229    /// ```
230    ///
231    /// # Errors
232    ///
233    /// Returns [`CpuContextError::InvalidThreadCount`] through
234    /// [`Error::Validation`] when `num_threads` is zero, a
235    /// [`Error::Validation`] error when the environment worker stack size is
236    /// malformed or too small, or [`Error::BackendSource`] when Rayon rejects
237    /// the thread pool.
238    pub fn with_threads(num_threads: usize) -> Result<Self> {
239        Self::with_threads_and_worker_stack(num_threads, worker_stack_bytes_from_env()?)
240    }
241
242    /// Create a CPU context with an explicit worker stack size.
243    ///
244    /// Provider calls run on pool workers, so the pool's stack bounds how deeply
245    /// a recursive provider implementation such as OpenBLAS's threaded LU can
246    /// descend; see [`DEFAULT_WORKER_STACK_BYTES`].
247    ///
248    /// # Examples
249    ///
250    /// ```
251    /// use tenferro_cpu::CpuContext;
252    ///
253    /// let ctx = CpuContext::with_threads_and_worker_stack(2, 32 << 20).unwrap();
254    /// assert_eq!(ctx.num_threads(), 2);
255    /// assert_eq!(ctx.worker_stack_bytes(), 32 << 20);
256    /// ```
257    ///
258    /// # Errors
259    ///
260    /// Returns [`Error::Validation`] when `num_threads` is zero or
261    /// `worker_stack_bytes` is below the minimum accepted stack, or
262    /// [`Error::BackendSource`] when Rayon rejects the thread pool.
263    pub fn with_threads_and_worker_stack(
264        num_threads: usize,
265        worker_stack_bytes: usize,
266    ) -> Result<Self> {
267        if num_threads == 0 {
268            return Err(Error::invalid_argument(
269                "CpuContext::with_threads",
270                "configuration",
271                "thread count must be at least 1",
272            ));
273        }
274        if worker_stack_bytes < MIN_WORKER_STACK_BYTES {
275            return Err(Error::invalid_argument(
276                "CpuContext::with_threads_and_worker_stack",
277                "configuration",
278                "worker stack size must be at least 65536 bytes",
279            ));
280        }
281        let execution_scope = Arc::new(ExecutionScopeState::default());
282        let pool = if num_threads == 1 {
283            None
284        } else {
285            let (startup_tx, startup_rx) = std::sync::mpsc::channel();
286            let worker_scope = Arc::clone(&execution_scope);
287            let pool = rayon::ThreadPoolBuilder::new()
288                .num_threads(num_threads)
289                .stack_size(worker_stack_bytes)
290                .start_handler(move |_| {
291                    register_worker_execution_scope(Arc::clone(&worker_scope));
292                    let _ = startup_tx.send(());
293                })
294                .build()
295                .map_err(|source| Error::backend_source("CpuContext::with_threads", source))?;
296            for _ in 0..num_threads {
297                startup_rx
298                    .recv()
299                    .map_err(|source| Error::backend_source("CpuContext::with_threads", source))?;
300            }
301            Some(Arc::new(pool))
302        };
303        Ok(Self {
304            num_threads,
305            worker_stack_bytes,
306            pool,
307            pinned_cpus: None,
308            execution_scope,
309            #[cfg(test)]
310            executor_install_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
311        })
312    }
313
314    /// Create a Rayon context whose workers are confined to an assigned CPU set.
315    ///
316    /// Every worker receives the whole `cpus` mask rather than one CPU, so threads
317    /// created by a provider inherit the full domain and keep their own
318    /// parallelism. The requested worker count cannot exceed the assigned CPU count.
319    ///
320    /// A real Rayon pool is constructed even when `num_threads` is one. The
321    /// worker count cannot exceed the assigned CPU count.
322    ///
323    /// # Examples
324    ///
325    /// ```
326    /// use tenferro_cpu::{process_cpu_affinity, CpuContext};
327    ///
328    /// if let Some(allowed) = process_cpu_affinity() {
329    ///     let one_cpu = tenferro_cpu::CpuSet::new([allowed.as_slice()[0]])?;
330    ///     let context = CpuContext::with_pinned_cpus(one_cpu.clone(), 1)?;
331    ///     assert_eq!(context.pinned_cpus(), Some(&one_cpu));
332    /// }
333    /// # Ok::<(), Box<dyn std::error::Error>>(())
334    /// ```
335    ///
336    /// # Errors
337    ///
338    /// Returns [`CpuContextError::InvalidThreadCount`] for zero workers,
339    /// [`CpuContextError::TooManyWorkers`] when the request exceeds the CPU
340    /// set, or an affinity error when workers cannot be confined.
341    pub fn with_pinned_cpus(
342        cpus: CpuSet,
343        num_threads: usize,
344    ) -> std::result::Result<Self, CpuContextError> {
345        let worker_stack_bytes = worker_stack_bytes_from_env()
346            .map_err(|source| CpuContextError::InvalidWorkerStack { source })?;
347        Self::with_pinned_cpus_and_worker_stack(
348            cpus,
349            num_threads,
350            worker_stack_bytes,
351            SystemThreadAffinity,
352        )
353    }
354
355    #[cfg(test)]
356    pub(crate) fn with_pinned_cpus_using<A: ThreadAffinity>(
357        cpus: CpuSet,
358        num_threads: usize,
359        affinity: A,
360    ) -> std::result::Result<Self, CpuContextError> {
361        let worker_stack_bytes = worker_stack_bytes_from_env()
362            .map_err(|source| CpuContextError::InvalidWorkerStack { source })?;
363        Self::with_pinned_cpus_and_worker_stack(cpus, num_threads, worker_stack_bytes, affinity)
364    }
365
366    pub(crate) fn with_pinned_cpus_and_worker_stack<A: ThreadAffinity>(
367        cpus: CpuSet,
368        num_threads: usize,
369        worker_stack_bytes: usize,
370        affinity: A,
371    ) -> std::result::Result<Self, CpuContextError> {
372        if worker_stack_bytes < MIN_WORKER_STACK_BYTES {
373            return Err(CpuContextError::InvalidWorkerStack {
374                source: Error::invalid_argument(
375                    "CpuContext::with_pinned_cpus",
376                    "configuration",
377                    "worker stack size must be at least 65536 bytes",
378                ),
379            });
380        }
381        if num_threads == 0 {
382            return Err(CpuContextError::InvalidThreadCount);
383        }
384        if num_threads > cpus.len() {
385            return Err(CpuContextError::TooManyWorkers {
386                workers: num_threads,
387                cpus: cpus.len(),
388            });
389        }
390
391        let execution_scope = Arc::new(ExecutionScopeState::default());
392        // Every worker is confined to the whole domain CPU set rather than to one
393        // CPU: threads created by a provider (BLAS/LAPACK) inherit the creating
394        // worker's mask, so a single-CPU mask would confine the provider's own
395        // thread team to one CPU and destroy its parallelism.
396        let domain_cpus = Arc::new(cpus.clone());
397        let (startup_tx, startup_rx) = std::sync::mpsc::channel();
398        let pool_domain_cpus = Arc::clone(&domain_cpus);
399        let worker_scope = Arc::clone(&execution_scope);
400        let pool = rayon::ThreadPoolBuilder::new()
401            .num_threads(num_threads)
402            .stack_size(worker_stack_bytes)
403            .spawn_handler(move |thread| {
404                let worker = thread.index();
405                let worker_cpus = Arc::clone(&pool_domain_cpus);
406                let startup_tx = startup_tx.clone();
407                let affinity = affinity.clone();
408                let worker_scope = Arc::clone(&worker_scope);
409                let mut builder =
410                    std::thread::Builder::new().name(format!("tenferro-cpu-{worker}"));
411                // Rayon applies its configured stack size only in its own spawn
412                // path, so a custom handler has to carry it to the OS thread.
413                if let Some(size) = thread.stack_size() {
414                    builder = builder.stack_size(size);
415                }
416                builder
417                    .spawn(move || {
418                        register_worker_execution_scope(Arc::clone(&worker_scope));
419                        let result = affinity.confine_current(&worker_cpus).and_then(|observed| {
420                            (observed.as_slice() == worker_cpus.as_slice())
421                                .then_some(())
422                                .ok_or_else(|| CpuAffinityError::Verification {
423                                    observed: observed.as_slice().to_vec(),
424                                })
425                        });
426                        let _ = startup_tx.send((worker, result));
427                        thread.run();
428                    })
429                    .map(|_| ())
430            })
431            .build()
432            .map_err(|source| CpuContextError::PoolBuild { source })?;
433        let pool = Arc::new(pool);
434        for _ in 0..num_threads {
435            let (worker, result) = startup_rx
436                .recv()
437                .map_err(|source| CpuContextError::WorkerStartupClosed { source })?;
438            if let Err(source) = result {
439                return Err(CpuContextError::WorkerAffinity { worker, source });
440            }
441        }
442        Ok(Self {
443            num_threads,
444            worker_stack_bytes,
445            pool: Some(pool),
446            pinned_cpus: Some(cpus),
447            execution_scope,
448            #[cfg(test)]
449            executor_install_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
450        })
451    }
452
453    fn single_threaded() -> Self {
454        Self {
455            num_threads: 1,
456            worker_stack_bytes: DEFAULT_WORKER_STACK_BYTES,
457            pool: None,
458            pinned_cpus: None,
459            execution_scope: Arc::new(ExecutionScopeState::default()),
460            #[cfg(test)]
461            executor_install_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
462        }
463    }
464
465    /// Return the stack size reserved for pool workers.
466    ///
467    /// A context that runs on the calling thread has no pool, so the reported
468    /// value is the size its workers would receive.
469    ///
470    /// # Examples
471    ///
472    /// ```
473    /// use tenferro_cpu::{CpuContext, DEFAULT_WORKER_STACK_BYTES};
474    ///
475    /// let ctx = CpuContext::with_threads(2).unwrap();
476    /// assert!(ctx.worker_stack_bytes() >= DEFAULT_WORKER_STACK_BYTES);
477    /// ```
478    pub fn worker_stack_bytes(&self) -> usize {
479        self.worker_stack_bytes
480    }
481
482    /// Return this context's CPU parallelism hint.
483    ///
484    /// # Examples
485    ///
486    /// ```
487    /// use tenferro_cpu::CpuContext;
488    ///
489    /// let ctx = CpuContext::with_threads(2).unwrap();
490    /// assert_eq!(ctx.num_threads(), 2);
491    /// ```
492    pub fn num_threads(&self) -> usize {
493        self.num_threads
494    }
495
496    /// Return the worker CPU domain for a pinned context.
497    ///
498    /// Legacy thread-count-only contexts return `None` because they do not own
499    /// worker affinity.
500    ///
501    /// # Examples
502    ///
503    /// ```
504    /// use tenferro_cpu::CpuContext;
505    ///
506    /// assert_eq!(CpuContext::with_threads(1)?.pinned_cpus(), None);
507    /// # Ok::<(), tenferro_tensor::Error>(())
508    /// ```
509    pub fn pinned_cpus(&self) -> Option<&CpuSet> {
510        self.pinned_cpus.as_ref()
511    }
512
513    /// Run a closure inside this context's CPU execution scope.
514    ///
515    /// # Examples
516    ///
517    /// ```
518    /// use tenferro_cpu::CpuContext;
519    ///
520    /// let ctx = CpuContext::with_threads(1).unwrap();
521    /// let value = ctx.install(|| 1 + 1);
522    /// assert_eq!(value, 2);
523    /// ```
524    pub fn install<R: Send>(&self, op: impl FnOnce() -> R + Send) -> R {
525        match &self.pool {
526            Some(pool) => pool.install(op),
527            None => op(),
528        }
529    }
530
531    pub(crate) fn install_if_needed<R: Send>(&self, op: impl FnOnce() -> R + Send) -> R {
532        if self.pool.is_some() && worker_execution_scope_matches(&self.execution_scope) {
533            op()
534        } else {
535            self.install(op)
536        }
537    }
538
539    #[cfg(test)]
540    pub(crate) fn owns_current_worker_for_test(&self) -> bool {
541        worker_execution_scope_matches(&self.execution_scope)
542    }
543
544    #[cfg(test)]
545    pub(crate) fn executor_install_calls_for_test(&self) -> usize {
546        self.executor_install_calls
547            .load(std::sync::atomic::Ordering::Relaxed)
548    }
549}
550
551impl CpuDomainExecutor for CpuContext {
552    fn capabilities(&self) -> CpuDomainExecutorCapabilities {
553        // INVARIANT: every CpuContext constructor rejects zero workers, and
554        // `num_threads` is private so it cannot be invalidated after creation.
555        let worker_count = match NonZeroUsize::new(self.num_threads) {
556            Some(worker_count) => worker_count,
557            None => unreachable!("CpuContext must contain at least one worker"),
558        };
559        CpuDomainExecutorCapabilities {
560            worker_count,
561            outer_parallelism: self.num_threads > 1,
562            inner_parallelism: if self.pool.is_some() {
563                CpuInnerParallelism::Rayon
564            } else {
565                CpuInnerParallelism::None
566            },
567            // This permits internal entry through the same executor. Public
568            // CpuBackend re-entry remains guarded by BACKEND_REENTRY_PANIC.
569            reentrancy: CpuExecutorReentrancy::SameExecutor,
570            affinity: if self.pinned_cpus.is_some() {
571                CpuExecutorAffinity::TenferroDomainVerified
572            } else {
573                CpuExecutorAffinity::None
574            },
575            shutdown: CpuExecutorShutdown::TenferroOwned,
576        }
577    }
578
579    fn submit(&self, jobs: &dyn ScopedCpuJobs) -> std::result::Result<(), CpuDomainExecutorError> {
580        let _scope = current_execution_owner().map(|owner| self.execution_scope.enter(owner));
581        if self.pool.is_none() {
582            return (0..jobs.len()).try_for_each(|index| jobs.run(index));
583        }
584        self.install_if_needed(|| {
585            (0..jobs.len())
586                .into_par_iter()
587                .try_for_each(|index| jobs.run(index))
588        })
589    }
590
591    fn install(
592        &self,
593        job: &mut dyn ScopedCpuJob,
594    ) -> std::result::Result<(), CpuDomainExecutorError> {
595        #[cfg(test)]
596        self.executor_install_calls
597            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
598        let _scope = current_execution_owner().map(|owner| self.execution_scope.enter(owner));
599        self.install_if_needed(|| job.run())
600    }
601
602    fn rayon_pool(&self) -> Option<&rayon::ThreadPool> {
603        self.pool.as_deref()
604    }
605}
606
607#[cfg(test)]
608mod tests;