Skip to main content

tensor4all_tensorbackend/
context.rs

1//! Explicit and optional process-global tenferro CPU execution contexts.
2
3use std::cell::Cell;
4use std::sync::{Arc, Mutex, OnceLock};
5
6use tenferro::{CompiledGraph, GraphCompiler, Runtime, Tensor, TracedGraph};
7use tenferro_ad::{AdContext, EagerRuntime};
8use tenferro_cpu::{BufferPoolStats, CpuBackend, CpuContext};
9use tenferro_tensor::{BackendSession, BackendSessionHost};
10
11/// Caller-owned execution domain used by context-aware tensor algorithms.
12///
13/// Values are validated against the exact runtime represented by the selected
14/// context; no implicit host/device transfer is performed by this enum.
15#[derive(Clone, Debug)]
16pub enum ExecutionContext {
17    /// Host execution through one caller-owned CPU context.
18    Cpu(Arc<CpuExecutionContext>),
19    /// CUDA execution through one caller-owned CUDA context.
20    #[cfg(feature = "tenferro-cuda")]
21    Cuda(Arc<crate::cuda::CudaExecutionContext>),
22}
23
24impl ExecutionContext {
25    /// Check whether this is the process-global CPU context.
26    ///
27    /// Compatibility boundary: legacy CPU-global callers route through the
28    /// historical host code paths (bitwise-identical numerics) while explicit
29    /// contexts use the scoped primitives. Without the global-defaults
30    /// feature there is no global context, so this always reports false.
31    pub fn is_global_default_cpu(&self) -> bool {
32        #[cfg(feature = "global-defaults")]
33        {
34            match self {
35                ExecutionContext::Cpu(context) => {
36                    let own = context.eager_runtime().map(|runtime| runtime.id());
37                    let global = defaults::default_eager_ctx().map(|runtime| runtime.id());
38                    matches!((own, global), (Ok(a), Ok(b)) if a == b)
39                }
40                #[cfg(feature = "tenferro-cuda")]
41                ExecutionContext::Cuda(_) => false,
42            }
43        }
44        #[cfg(not(feature = "global-defaults"))]
45        {
46            let _ = self;
47            false
48        }
49    }
50}
51
52/// Error returned by explicit CPU context graph or eager-runtime operations.
53///
54/// The original tenferro diagnostic is retained as the error source.
55///
56/// # Examples
57///
58/// ```
59/// use std::error::Error;
60/// use std::sync::Arc;
61/// use tensor4all_tensorbackend::CpuExecutionContextError;
62///
63/// let error = CpuExecutionContextError::Initialization {
64///     component: "graph runtime",
65///     source: Arc::new(std::io::Error::other("registration failed")),
66/// };
67/// assert!(error.source().is_some());
68/// ```
69#[derive(Debug, Clone, thiserror::Error)]
70pub enum CpuExecutionContextError {
71    /// A context-owned graph or eager runtime could not be initialized.
72    #[error("failed to initialize {component}: {source}")]
73    Initialization {
74        /// Context component being initialized.
75        component: &'static str,
76        /// Original tenferro diagnostic.
77        #[source]
78        source: Arc<dyn std::error::Error + Send + Sync + 'static>,
79    },
80    /// Graph compilation or execution failed.
81    #[error("CPU graph {operation} failed: {source}")]
82    Graph {
83        /// Graph operation that failed.
84        operation: &'static str,
85        /// Original tenferro diagnostic.
86        #[source]
87        source: Arc<dyn std::error::Error + Send + Sync + 'static>,
88    },
89}
90
91const CANONICAL_SESSION_REENTRY_MESSAGE: &str = "recursive tensorbackend canonical session entry";
92
93thread_local! {
94    static CANONICAL_SESSION_ACTIVE: Cell<bool> = const { Cell::new(false) };
95}
96
97struct CanonicalSessionGuard {
98    previous: bool,
99}
100
101impl CanonicalSessionGuard {
102    fn assert_inactive() {
103        CANONICAL_SESSION_ACTIVE.with(|active| {
104            assert!(!active.get(), "{CANONICAL_SESSION_REENTRY_MESSAGE}");
105        });
106    }
107
108    fn enter() -> Self {
109        Self::assert_inactive();
110        CANONICAL_SESSION_ACTIVE.with(|active| Self {
111            previous: active.replace(true),
112        })
113    }
114}
115
116impl Drop for CanonicalSessionGuard {
117    fn drop(&mut self) {
118        CANONICAL_SESSION_ACTIVE.with(|active| active.set(self.previous));
119    }
120}
121
122/// Run one concrete session on `backend` under the canonical session guard.
123fn run_canonical_session<R: Send>(
124    backend: &mut CpuBackend,
125    f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
126) -> R {
127    backend.with_backend_session(|session| {
128        let _guard = CanonicalSessionGuard::enter();
129        f(session)
130    })
131}
132
133impl CpuExecutionContextError {
134    fn initialization(
135        component: &'static str,
136        source: impl std::error::Error + Send + Sync + 'static,
137    ) -> Self {
138        Self::Initialization {
139            component,
140            source: Arc::new(source),
141        }
142    }
143
144    fn graph(
145        operation: &'static str,
146        source: impl std::error::Error + Send + Sync + 'static,
147    ) -> Self {
148        Self::Graph {
149            operation,
150            source: Arc::new(source),
151        }
152    }
153}
154
155struct GraphState {
156    compiler: GraphCompiler,
157    runtime: Runtime,
158    backend: CpuBackend,
159}
160
161/// Caller-owned CPU execution domain for plain, graph, and eager-AD work.
162///
163/// The supplied backend is the only source of CPU execution resources for every
164/// entry from a thread that is not a Rayon worker. A plain session entered from a
165/// Rayon worker runs inline and single-threaded on a context-local pool-less CPU
166/// backend instead: a worker that waits for a pool install is handed more of the
167/// enclosing pool's work, and that work may enter a session itself
168/// (tensor4all-rs#830). Backend clones preserve its runtime identity; this
169/// constructor never uses `CpuBackend::new`, `CpuContext::from_env`, or a
170/// process-global fallback.
171/// Graph preparation caches and the eager runtime are owned by this context and
172/// are released when it is dropped.
173///
174/// # Examples
175///
176/// ```
177/// use tensor4all_tensorbackend::CpuExecutionContext;
178/// use tenferro_cpu::CpuBackend;
179///
180/// let context = CpuExecutionContext::from_backend(CpuBackend::with_threads(1)?);
181/// let threads = context.with_backend(|backend| backend.num_threads());
182/// assert_eq!(threads, 1);
183/// # Ok::<(), Box<dyn std::error::Error>>(())
184/// ```
185pub struct CpuExecutionContext {
186    backend: Mutex<CpuBackend>,
187    inline: OnceLock<CpuBackend>,
188    graph: OnceLock<Result<Mutex<GraphState>, CpuExecutionContextError>>,
189    eager: OnceLock<Result<Arc<EagerRuntime>, CpuExecutionContextError>>,
190}
191
192impl std::fmt::Debug for CpuExecutionContext {
193    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
194        f.debug_struct("CpuExecutionContext")
195            .field("graph_initialized", &self.graph.get().is_some())
196            .field("eager_initialized", &self.eager.get().is_some())
197            .finish_non_exhaustive()
198    }
199}
200
201impl CpuExecutionContext {
202    /// Create an execution context from a caller-selected CPU backend.
203    ///
204    /// Runtime construction is lazy, so creating a context cannot fail and does
205    /// not allocate another executor or consult environment configuration.
206    pub fn from_backend(backend: CpuBackend) -> Self {
207        Self {
208            backend: Mutex::new(backend),
209            inline: OnceLock::new(),
210            graph: OnceLock::new(),
211            eager: OnceLock::new(),
212        }
213    }
214
215    /// Run a plain tensor operation with this context's backend.
216    ///
217    /// The closure runs while the context-local backend lock is held. A poisoned
218    /// lock is recovered because tenferro validates every new backend session.
219    pub fn with_backend<R>(&self, f: impl FnOnce(&mut CpuBackend) -> R) -> R {
220        let mut backend = match self.backend.lock() {
221            Ok(guard) => guard,
222            Err(poisoned) => poisoned.into_inner(),
223        };
224        f(&mut backend)
225    }
226
227    pub(crate) fn with_session<R: Send>(
228        &self,
229        f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
230    ) -> R {
231        CanonicalSessionGuard::assert_inactive();
232        // A Rayon worker can be handed more of the enclosing pool's work while
233        // tenferro installs this session into the context's own pool. That stolen
234        // work may enter a session itself, on a thread whose tenferro execution is
235        // already active, which tenferro rejects. Entering from a worker therefore
236        // runs the session inline and single-threaded on that worker and leaves
237        // parallelism to the enclosing pool.
238        if rayon::current_thread_index().is_some() {
239            let mut backend = self.inline_backend();
240            return run_canonical_session(&mut backend, f);
241        }
242        let mut backend = match self.backend.lock() {
243            Ok(guard) => guard,
244            Err(poisoned) => poisoned.into_inner(),
245        };
246        run_canonical_session(&mut backend, f)
247    }
248
249    /// Backend used for one session entered from a Rayon worker.
250    ///
251    /// `CpuContext::with_threads(1)` owns no Rayon pool, so its executor runs
252    /// every operation on the calling thread and never installs a session into a
253    /// pool the caller does not belong to.
254    fn inline_backend(&self) -> CpuBackend {
255        self.inline
256            .get_or_init(|| match CpuContext::with_threads(1) {
257                Ok(context) => CpuBackend::from_context(Arc::new(context)),
258                // INVARIANT: `CpuContext::with_threads` rejects only a zero worker
259                // count, which this call never passes, so this arm is unreachable.
260                // Keeping the supplied backend is a last resort rather than a
261                // policy: it is the only remaining backend this context owns.
262                Err(_) => self.backend_clone(),
263            })
264            .clone()
265    }
266
267    fn backend_clone(&self) -> CpuBackend {
268        self.with_backend(|backend| backend.clone())
269    }
270
271    fn graph_state(&self) -> Result<&Mutex<GraphState>, CpuExecutionContextError> {
272        self.graph
273            .get_or_init(|| {
274                let backend = self.backend_clone();
275                build_graph_runtime(&backend).map(|runtime| {
276                    Mutex::new(GraphState {
277                        compiler: GraphCompiler::new(),
278                        runtime,
279                        backend,
280                    })
281                })
282            })
283            .as_ref()
284            .map_err(Clone::clone)
285    }
286
287    fn with_graph_state<R>(
288        &self,
289        f: impl FnOnce(&mut GraphCompiler, &mut Runtime, &mut CpuBackend) -> R,
290    ) -> Result<R, CpuExecutionContextError> {
291        let mut graph = match self.graph_state()?.lock() {
292            Ok(guard) => guard,
293            Err(poisoned) => poisoned.into_inner(),
294        };
295        let GraphState {
296            compiler,
297            runtime,
298            backend,
299        } = &mut *graph;
300        Ok(f(compiler, runtime, backend))
301    }
302
303    /// Compile a backend-neutral traced graph using this context's compiler cache.
304    ///
305    /// # Errors
306    ///
307    /// Returns [`CpuExecutionContextError`] when graph-runtime initialization or
308    /// graph compilation fails.
309    pub fn compile_graph(
310        &self,
311        graph: &TracedGraph,
312    ) -> Result<CompiledGraph, CpuExecutionContextError> {
313        self.with_graph_state(|compiler, _, _| compiler.compile_traced_graph(graph))?
314            .map_err(|source| CpuExecutionContextError::graph("compilation", source))
315    }
316
317    /// Execute a compiled graph in this context's runtime and prepared-plan cache.
318    ///
319    /// `CompiledGraph` is backend-neutral. Backend-prepared executables and
320    /// workspaces never leave this context-owned runtime.
321    ///
322    /// # Errors
323    ///
324    /// Returns [`CpuExecutionContextError`] when runtime initialization,
325    /// preparation, or execution fails.
326    pub fn run_graph(
327        &self,
328        graph: &CompiledGraph,
329        inputs: &[&Tensor],
330    ) -> Result<Vec<Tensor>, CpuExecutionContextError> {
331        self.with_graph_state(|_, runtime, _| runtime.run_compiled(graph, inputs))?
332            .map_err(|source| CpuExecutionContextError::graph("execution", source))
333    }
334
335    /// Return this context's eager reverse-AD runtime.
336    ///
337    /// Repeated calls return the same runtime and therefore the same eager
338    /// compilation cache.
339    ///
340    /// # Errors
341    ///
342    /// Returns [`CpuExecutionContextError`] when linalg AD-rule or eager-runtime
343    /// registration fails.
344    pub fn eager_runtime(&self) -> Result<Arc<EagerRuntime>, CpuExecutionContextError> {
345        self.eager
346            .get_or_init(|| build_eager_runtime(self.backend_clone()))
347            .as_ref()
348            .map(Arc::clone)
349            .map_err(Clone::clone)
350    }
351
352    /// Return statistics for this context's runtime-owned graph caches.
353    ///
354    /// # Errors
355    ///
356    /// Returns [`CpuExecutionContextError`] when graph initialization or the
357    /// cache statistics query fails.
358    pub fn graph_cache_stats(
359        &self,
360    ) -> Result<tenferro::RuntimeCacheStats, CpuExecutionContextError> {
361        self.with_graph_state(|_, runtime, _| runtime.cache_stats())?
362            .map_err(|source| CpuExecutionContextError::graph("cache statistics", source))
363    }
364
365    /// Return retained-buffer statistics for this context's graph backend.
366    ///
367    /// # Errors
368    ///
369    /// Returns [`CpuExecutionContextError`] when graph initialization or the
370    /// backend statistics query fails.
371    pub fn graph_buffer_pool_stats(&self) -> Result<BufferPoolStats, CpuExecutionContextError> {
372        self.with_graph_state(|_, _, backend| backend.buffer_pool_stats())?
373            .map_err(|source| CpuExecutionContextError::graph("buffer-pool statistics", source))
374    }
375
376    /// Release retained buffers owned by this context's graph backend.
377    ///
378    /// # Errors
379    ///
380    /// Returns [`CpuExecutionContextError`] when graph initialization or reset
381    /// fails.
382    pub fn reset_graph_buffer_pool(&self) -> Result<(), CpuExecutionContextError> {
383        self.with_graph_state(|_, _, backend| backend.reset_buffer_pool())?
384            .map_err(|source| CpuExecutionContextError::graph("buffer-pool reset", source))
385    }
386
387    /// Recreate this context's graph runtime and release its prepared caches.
388    ///
389    /// # Errors
390    ///
391    /// Returns [`CpuExecutionContextError`] when runtime reconstruction or
392    /// buffer release fails.
393    pub fn reset_graph_runtime(&self) -> Result<(), CpuExecutionContextError> {
394        self.with_graph_state(|compiler, runtime, backend| {
395            let replacement = build_graph_runtime(backend)?;
396            *compiler = GraphCompiler::new();
397            let old = std::mem::replace(runtime, replacement);
398            drop(old);
399            backend
400                .reset_buffer_pool()
401                .map_err(|source| CpuExecutionContextError::graph("buffer-pool reset", source))
402        })??;
403        Ok(())
404    }
405}
406
407fn build_graph_runtime(backend: &CpuBackend) -> Result<Runtime, CpuExecutionContextError> {
408    let mut builder = Runtime::builder();
409    builder
410        .register_engine(
411            tenferro_cpu::runtime_engine_registration(backend).map_err(|source| {
412                CpuExecutionContextError::initialization("graph CPU engine", source)
413            })?,
414        )
415        .map_err(|source| CpuExecutionContextError::initialization("graph CPU engine", source))?;
416    builder
417        .install_extension_module(
418            tenferro_einsum::extension_module::<CpuBackend>(
419                tenferro_cpu::runtime_engine_id().map_err(|source| {
420                    CpuExecutionContextError::initialization("einsum extension", source)
421                })?,
422            )
423            .map_err(|source| {
424                CpuExecutionContextError::initialization("einsum extension", source)
425            })?,
426        )
427        .map_err(|source| CpuExecutionContextError::initialization("einsum extension", source))?;
428    builder
429        .build()
430        .map_err(|source| CpuExecutionContextError::initialization("graph runtime", source))
431}
432
433fn build_eager_runtime(backend: CpuBackend) -> Result<Arc<EagerRuntime>, CpuExecutionContextError> {
434    let ad_context = AdContext::builder()
435        .with_semantic_extension_rules(tenferro_linalg::semantic_ad_rules().map_err(|source| {
436            CpuExecutionContextError::initialization("linalg AD rules", source)
437        })?)
438        .map_err(|source| CpuExecutionContextError::initialization("linalg AD rules", source))?
439        .build()
440        .map_err(|source| CpuExecutionContextError::initialization("AD context", source))?;
441    let runtime = EagerRuntime::with_cpu_backend_and_ad_context(backend, &ad_context)
442        .map_err(|source| CpuExecutionContextError::initialization("eager runtime", source))?;
443    // [AI Supplied] Install the built-in extension modules before publishing
444    // the shared context. Lazy first use reconfigures the runtime and advances
445    // its epoch; doing that after tensors have prepared AD derivatives can
446    // invalidate those prepared programs under parallel first use.
447    let engine_id = tenferro_cpu::runtime_engine_id()
448        .map_err(|source| CpuExecutionContextError::initialization("CPU runtime engine", source))?;
449    let einsum_module = tenferro_einsum::extension_module::<CpuBackend>(engine_id.clone())
450        .map_err(|source| CpuExecutionContextError::initialization("einsum extension", source))?;
451    runtime
452        .install_extension_module(einsum_module)
453        .map_err(|source| CpuExecutionContextError::initialization("einsum runtime", source))?;
454    let linalg_module = tenferro_linalg::extension_module::<CpuBackend>(engine_id)
455        .map_err(|source| CpuExecutionContextError::initialization("linalg extension", source))?;
456    runtime
457        .install_extension_module(linalg_module)
458        .map_err(|source| CpuExecutionContextError::initialization("linalg runtime", source))?;
459    Ok(runtime)
460}
461
462#[cfg(feature = "global-defaults")]
463mod defaults {
464    use super::*;
465    use tenferro_cpu::CpuContext;
466
467    static DEFAULT_CONTEXT: OnceLock<Arc<CpuExecutionContext>> = OnceLock::new();
468
469    #[cfg(test)]
470    thread_local! {
471        static FORCE_EAGER_CONTEXT_FAILURE: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
472    }
473
474    #[cfg(test)]
475    static DEFAULT_CONTEXT_HITS: std::sync::atomic::AtomicUsize =
476        std::sync::atomic::AtomicUsize::new(0);
477
478    fn default_context() -> &'static Arc<CpuExecutionContext> {
479        DEFAULT_CONTEXT.get_or_init(|| {
480            #[cfg(test)]
481            DEFAULT_CONTEXT_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
482            Arc::new(CpuExecutionContext::from_backend(CpuBackend::from_context(
483                Arc::new(CpuContext::from_env()),
484            )))
485        })
486    }
487
488    /// Error returned when the process-global eager AD runtime cannot be initialized.
489    ///
490    /// # Examples
491    ///
492    /// ```
493    /// use std::error::Error;
494    /// use std::sync::Arc;
495    /// use tensor4all_tensorbackend::EagerContextError;
496    ///
497    /// let error = EagerContextError::Registration {
498    ///     source: Arc::new(std::io::Error::other("registration failed")),
499    /// };
500    /// assert!(error.source().is_some());
501    /// ```
502    #[derive(Debug, Clone, thiserror::Error)]
503    pub enum EagerContextError {
504        /// The tenferro linalg AD extension rule could not be registered.
505        #[error("failed to register tenferro linalg AD rule: {source}")]
506        Registration {
507            /// Original diagnostic returned by tenferro.
508            #[source]
509            source: Arc<dyn std::error::Error + Send + Sync + 'static>,
510        },
511    }
512
513    /// Run a closure against the optional process-global CPU backend.
514    pub fn with_default_backend<R>(f: impl FnOnce(&mut CpuBackend) -> R) -> R {
515        default_context().with_backend(f)
516    }
517
518    pub(crate) fn with_default_session<R: Send>(
519        f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
520    ) -> R {
521        default_context().with_session(f)
522    }
523
524    pub(crate) fn with_default_graph_runtime<R>(
525        f: impl FnOnce(&mut GraphCompiler, &Runtime, &mut CpuBackend) -> R,
526    ) -> anyhow::Result<R> {
527        default_context()
528            .with_graph_state(|compiler, runtime, backend| f(compiler, runtime, backend))
529            .map_err(anyhow::Error::new)
530    }
531
532    pub(crate) fn default_engine_buffer_pool_stats() -> anyhow::Result<BufferPoolStats> {
533        default_context()
534            .graph_buffer_pool_stats()
535            .map_err(anyhow::Error::new)
536    }
537
538    pub(crate) fn reset_default_engine_buffer_pool() -> anyhow::Result<()> {
539        default_context()
540            .reset_graph_buffer_pool()
541            .map_err(anyhow::Error::new)
542    }
543
544    pub(crate) fn reset_default_engine() -> anyhow::Result<()> {
545        default_context()
546            .reset_graph_runtime()
547            .map_err(anyhow::Error::new)
548    }
549
550    /// Return the optional process-global eager context used by convenience APIs.
551    ///
552    /// # Errors
553    ///
554    /// Returns [`EagerContextError::Registration`] when eager runtime
555    /// initialization fails.
556    ///
557    /// # Examples
558    ///
559    /// ```
560    /// use std::sync::Arc;
561    /// use tensor4all_tensorbackend::default_eager_ctx;
562    ///
563    /// let first = default_eager_ctx().unwrap();
564    /// let second = default_eager_ctx().unwrap();
565    /// assert!(Arc::ptr_eq(&first, &second));
566    /// ```
567    pub fn default_eager_ctx() -> Result<Arc<EagerRuntime>, EagerContextError> {
568        #[cfg(test)]
569        if FORCE_EAGER_CONTEXT_FAILURE.with(std::cell::Cell::get) {
570            return Err(EagerContextError::Registration {
571                source: Arc::new(std::io::Error::other(
572                    "forced default eager context registration failure",
573                )),
574            });
575        }
576        default_context()
577            .eager_runtime()
578            .map_err(|source| EagerContextError::Registration {
579                source: Arc::new(source),
580            })
581    }
582
583    /// Borrow the process-global CPU execution context.
584    ///
585    /// Compatibility entry for CPU-global convenience APIs (e.g. the legacy
586    /// context-free SRC entry): host tensors constructed through the global
587    /// default belong to this exact context, so they validate against it.
588    /// New code must take a caller-owned context instead of consulting this.
589    ///
590    /// # Examples
591    ///
592    /// ```
593    /// use tensor4all_tensorbackend::{default_cpu_execution_context, ExecutionContext};
594    ///
595    /// let context = ExecutionContext::Cpu(default_cpu_execution_context());
596    /// assert!(matches!(context, ExecutionContext::Cpu(_)));
597    /// ```
598    pub fn default_cpu_execution_context() -> Arc<CpuExecutionContext> {
599        Arc::clone(default_context())
600    }
601
602    #[cfg(test)]
603    pub(crate) fn default_context_hits() -> usize {
604        DEFAULT_CONTEXT_HITS.load(std::sync::atomic::Ordering::Relaxed)
605    }
606
607    #[cfg(test)]
608    pub(crate) fn with_forced_eager_context_failure<T>(f: impl FnOnce() -> T) -> T {
609        let previous = FORCE_EAGER_CONTEXT_FAILURE.with(|failure| failure.replace(true));
610        let result = f();
611        FORCE_EAGER_CONTEXT_FAILURE.with(|failure| failure.set(previous));
612        result
613    }
614}
615
616#[cfg(all(test, feature = "global-defaults"))]
617pub(crate) use defaults::with_forced_eager_context_failure;
618#[cfg(feature = "global-defaults")]
619pub use defaults::{
620    default_cpu_execution_context, default_eager_ctx, with_default_backend, EagerContextError,
621};
622#[cfg(feature = "global-defaults")]
623pub(crate) use defaults::{
624    default_engine_buffer_pool_stats, reset_default_engine, reset_default_engine_buffer_pool,
625    with_default_graph_runtime, with_default_session,
626};
627
628#[cfg(test)]
629mod tests {
630    use std::num::NonZeroUsize;
631    use std::sync::mpsc;
632    use std::time::Duration;
633
634    use super::*;
635    use tenferro::program::{CoreSemanticOp, ProgramInputSpec};
636    use tenferro::{DType, TensorSessionOpsExt, TraceContext};
637    use tenferro_ad::EagerTensor;
638    use tenferro_cpu::{CpuContext, ExternalCpuDomain};
639    use tenferro_tensor::CpuDomainId;
640
641    fn context() -> CpuExecutionContext {
642        CpuExecutionContext::from_backend(CpuBackend::with_threads(1).unwrap())
643    }
644
645    #[test]
646    fn explicit_session_runs_a_concrete_operation() {
647        let context = context();
648        let lhs = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0]).unwrap();
649        let rhs = Tensor::from_vec_col_major(vec![2, 1], vec![5.0_f64, 6.0]).unwrap();
650        let result = context
651            .with_session(|session| lhs.matmul(&rhs, session))
652            .unwrap();
653
654        assert_eq!(result.as_slice::<f64>().unwrap(), &[23.0, 34.0]);
655    }
656
657    #[test]
658    fn recursive_session_entry_fails_before_lock_and_restores_guard() {
659        let context = context();
660        let panic = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
661            context.with_session(|_| context.with_session(|_| ()))
662        }))
663        .expect_err("recursive canonical session entry should panic");
664        let message = panic
665            .downcast_ref::<&str>()
666            .copied()
667            .or_else(|| panic.downcast_ref::<String>().map(String::as_str))
668            .expect("recursive entry panic should contain a string message");
669        assert_eq!(message, CANONICAL_SESSION_REENTRY_MESSAGE);
670
671        assert_eq!(context.with_session(|_| 7usize), 7);
672    }
673
674    #[test]
675    fn explicit_plain_graph_and_eager_paths_share_only_the_supplied_backend() {
676        let context = context();
677        assert!(format!("{context:?}").contains("graph_initialized: false"));
678        assert_eq!(context.with_backend(|backend| backend.num_threads()), 1);
679
680        let mut trace = TraceContext::new();
681        let input = trace
682            .input(ProgramInputSpec::new(DType::F64, [2_usize.into()]))
683            .unwrap();
684        let output = trace.add_op(CoreSemanticOp::Neg, &[input]).unwrap()[0];
685        let graph = trace.finish(&[output]).unwrap();
686        let compiled = context.compile_graph(&graph).unwrap();
687        let input = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, -2.0]).unwrap();
688        let output = context.run_graph(&compiled, &[&input]).unwrap();
689        assert_eq!(output[0].as_slice::<f64>().unwrap(), &[-1.0, 2.0]);
690        context.run_graph(&compiled, &[&input]).unwrap();
691        let cached = context.graph_cache_stats().unwrap().prepared_plans;
692        assert!(cached.entries > 0);
693        assert!(cached.hits > 0);
694        context.reset_graph_runtime().unwrap();
695        assert_eq!(
696            context.graph_cache_stats().unwrap().prepared_plans.entries,
697            0
698        );
699
700        let eager = context.eager_runtime().unwrap();
701        assert!(Arc::ptr_eq(&eager, &context.eager_runtime().unwrap()));
702    }
703
704    #[test]
705    fn separate_eager_contexts_reject_cross_context_operations() {
706        let first = context().eager_runtime().unwrap();
707        let second = context().eager_runtime().unwrap();
708        let a = EagerTensor::from_tensor_in(
709            Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(),
710            first,
711        )
712        .unwrap();
713        let b = EagerTensor::from_tensor_in(
714            Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap(),
715            second,
716        )
717        .unwrap();
718        assert!(matches!(
719            a.add(&b),
720            Err(tenferro_ad::Error::ContextMismatch { .. })
721        ));
722    }
723
724    #[test]
725    fn caller_managed_backend_remains_caller_owned_after_context_drop() {
726        let executor = Arc::new(CpuContext::with_threads(1).unwrap());
727        let id = CpuDomainId::new(7);
728        let domain =
729            ExternalCpuDomain::new_caller_managed(id, executor.clone(), NonZeroUsize::MIN).unwrap();
730        let backend = CpuBackend::from_external_managed_domains(id, [domain]).unwrap();
731        let context = CpuExecutionContext::from_backend(backend);
732        assert_eq!(context.with_backend(|backend| backend.num_threads()), 1);
733        drop(context);
734        assert_eq!(executor.num_threads(), 1);
735    }
736
737    #[test]
738    fn independent_contexts_do_not_share_a_backend_mutex() {
739        let first = Arc::new(context());
740        let second = Arc::new(context());
741        let (entered_tx, entered_rx) = mpsc::channel();
742        let (release_tx, release_rx) = mpsc::channel();
743        let release_rx = Arc::new(Mutex::new(release_rx));
744        let handles = [first, second].map(|context| {
745            let entered_tx = entered_tx.clone();
746            let release_rx = Arc::clone(&release_rx);
747            std::thread::spawn(move || {
748                context.with_backend(|_| {
749                    entered_tx.send(()).unwrap();
750                    release_rx.lock().unwrap().recv().unwrap();
751                });
752            })
753        });
754        entered_rx.recv_timeout(Duration::from_secs(2)).unwrap();
755        entered_rx.recv_timeout(Duration::from_secs(2)).unwrap();
756        release_tx.send(()).unwrap();
757        release_tx.send(()).unwrap();
758        for handle in handles {
759            handle.join().unwrap();
760        }
761    }
762
763    #[test]
764    fn session_from_a_rayon_worker_completes_without_waiting_on_the_context_pool() {
765        use rayon::prelude::*;
766
767        // The context owns a two-worker Rayon pool, and the enclosing pool hands
768        // the worker waiting for a session install more of its own items. A
769        // session entry that installs into the context pool therefore re-enters a
770        // session on a thread whose execution is already active; the one-worker
771        // case makes that sequence certain.
772        let context = Arc::new(CpuExecutionContext::from_backend(
773            CpuBackend::with_threads(2).unwrap(),
774        ));
775        for enclosing_workers in [1usize, 2] {
776            let pool = Arc::new(
777                rayon::ThreadPoolBuilder::new()
778                    .num_threads(enclosing_workers)
779                    .build()
780                    .unwrap(),
781            );
782            let (sender, receiver) = mpsc::channel();
783            let context = Arc::clone(&context);
784            std::thread::spawn(move || {
785                let results = pool.install(|| {
786                    (0..2usize)
787                        .into_par_iter()
788                        .map(|_| {
789                            let lhs = Tensor::from_vec_col_major(
790                                vec![2, 2],
791                                vec![1.0_f64, 2.0, 3.0, 4.0],
792                            )
793                            .unwrap();
794                            let rhs =
795                                Tensor::from_vec_col_major(vec![2, 1], vec![5.0_f64, 6.0]).unwrap();
796                            context
797                                .with_session(|session| lhs.matmul(&rhs, session))
798                                .unwrap()
799                                .as_slice::<f64>()
800                                .unwrap()
801                                .to_vec()
802                        })
803                        .collect::<Vec<_>>()
804                });
805                let _ = sender.send(results);
806            });
807
808            let results = receiver.recv_timeout(Duration::from_secs(60)).expect(
809                "a session entered from a Rayon worker must finish instead of waiting on the context pool",
810            );
811            assert_eq!(results.len(), 2, "enclosing workers = {enclosing_workers}");
812            for values in results {
813                assert_eq!(values, vec![23.0, 34.0]);
814            }
815        }
816    }
817
818    #[cfg(feature = "global-defaults")]
819    #[test]
820    fn explicit_paths_do_not_initialize_the_default_context() {
821        let before = defaults::default_context_hits();
822        let context = context();
823        context.with_backend(|backend| assert_eq!(backend.num_threads(), 1));
824        context.eager_runtime().unwrap();
825        assert_eq!(defaults::default_context_hits(), before);
826    }
827}