Skip to main content

tenferro_ops/ad/
context.rs

1//! AD context for guard-based shape resolution and value metadata queries.
2//!
3//! During AD graph construction, linalg rules such as SVD, QR, and LU need
4//! concrete matrix dimensions to choose between structurally different
5//! subgraphs. `ShapeGuardContext` records those dimension comparisons as guards
6//! so cached AD graphs can later be invalidated when the observed shape
7//! relationship changes.
8
9use std::cmp::Ordering;
10use std::collections::HashMap;
11use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering};
12#[cfg(feature = "autodiff")]
13use std::sync::Arc;
14use std::sync::{Mutex, OnceLock};
15
16use computegraph::graph::Graph;
17use computegraph::types::{ValueKey, ValueRef};
18use tenferro_tensor::DType;
19
20#[cfg(feature = "autodiff")]
21use crate::ad::{ADRuleError, ADRuleKind, ExtensionAdDispatcher};
22use crate::dim_expr::{DimExpr, DimExprEvalError};
23use crate::shape_extent::ShapeExtent;
24use crate::std_tensor_op::StdTensorOp;
25use crate::sym_dim::SymDim;
26
27type MetadataMap = HashMap<ValueKey<StdTensorOp>, TensorMeta>;
28
29type GlobalMetadataMap = HashMap<ValueKey<StdTensorOp>, GlobalMetadataEntry>;
30
31#[derive(Clone, Debug)]
32struct GlobalMetadataEntry {
33    stack: Vec<GlobalMetadataRegistration>,
34}
35
36#[derive(Clone, Debug)]
37struct GlobalMetadataRegistration {
38    token: u64,
39    meta: TensorMeta,
40}
41
42#[derive(Clone, Debug)]
43struct ScopedGlobalMetadataRegistration {
44    key: ValueKey<StdTensorOp>,
45    token: u64,
46}
47
48/// Error returned when the process-global AD metadata registry is unavailable.
49#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
50pub enum MetadataRegistryError {
51    /// A previous panic poisoned the global metadata mutex.
52    #[error("AD global metadata registry lock poisoned")]
53    LockPoisoned,
54}
55
56/// Error returned when shape-guard metadata cannot be resolved.
57///
58/// # Examples
59///
60/// ```
61/// use tenferro_ops::ShapeGuardError;
62///
63/// let error = ShapeGuardError::LocalWithoutAttachedGraph { local_id: 0 };
64/// assert!(error.to_string().contains("attached graph"));
65/// ```
66#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
67pub enum ShapeGuardError {
68    /// A local graph value was queried before a graph was attached.
69    #[error("cannot resolve local value {local_id} without an attached graph")]
70    LocalWithoutAttachedGraph {
71        /// Graph-local value id.
72        local_id: usize,
73    },
74    /// A local graph value id is outside the attached graph's value table.
75    #[error("local value {local_id} is out of bounds for the attached graph")]
76    LocalOutOfBounds {
77        /// Graph-local value id.
78        local_id: usize,
79    },
80    /// No metadata was registered for the resolved value key.
81    #[error("missing TensorMeta for {key:?}")]
82    MissingMetadata {
83        /// Resolved value key.
84        key: ValueKey<StdTensorOp>,
85    },
86    /// Metadata exists, but at least one axis is only bounded or unknown.
87    #[error("TensorMeta for {key:?} does not have an exact shape; query extents instead")]
88    NonExactShape {
89        /// Resolved value key.
90        key: ValueKey<StdTensorOp>,
91    },
92}
93
94/// Result type used by shape-guard metadata queries.
95///
96/// The error side preserves a [`ShapeGuardFailure`] wrapper so an AD callback
97/// can retain the original [`ShapeGuardError`] even when a foreign callback
98/// protocol accepts only a rendered message.
99///
100/// # Examples
101///
102/// ```
103/// use tenferro_ops::ShapeGuardResult;
104///
105/// let result: ShapeGuardResult<()> = Ok(());
106/// assert!(result.is_ok());
107/// ```
108pub type ShapeGuardResult<T> = Result<T, ShapeGuardFailure>;
109
110#[cfg(feature = "autodiff")]
111impl From<ShapeGuardFailure> for ADRuleError {
112    fn from(err: ShapeGuardFailure) -> Self {
113        err.record_for_ad_boundary();
114        ADRuleError::invalid_input(
115            "tenferro.shape_guard",
116            ADRuleKind::Jvp,
117            err.typed_source().to_string(),
118        )
119    }
120}
121
122/// Error returned by a shape-guard metadata query.
123///
124/// The public [`ShapeGuardError`] remains the typed source. The private side
125/// channel is shared with the owning [`ShapeGuardContext`] so an external
126/// message-only AD callback can report the same typed source at the runtime
127/// boundary without changing the callback protocol.
128///
129/// # Examples
130///
131/// ```
132/// use computegraph::types::{ValueKey, ValueRef};
133/// use tenferro_ops::input_key::TensorInputKey;
134/// use tenferro_ops::std_tensor_op::StdTensorOp;
135/// use tenferro_ops::{ShapeGuardContext, ShapeGuardError};
136///
137/// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 8 });
138/// let value = ValueRef::External(key);
139/// let mut ctx = ShapeGuardContext::default();
140/// let failure = ctx.shape_of(&value).unwrap_err();
141/// assert!(matches!(
142///     failure.typed_source(),
143///     ShapeGuardError::MissingMetadata { .. }
144/// ));
145/// ```
146#[derive(Clone, Debug)]
147pub struct ShapeGuardFailure {
148    source: ShapeGuardError,
149    #[cfg(feature = "autodiff")]
150    deferred: Arc<Mutex<Option<ShapeGuardError>>>,
151}
152
153impl ShapeGuardFailure {
154    #[cfg(feature = "autodiff")]
155    fn new(source: ShapeGuardError, deferred: Arc<Mutex<Option<ShapeGuardError>>>) -> Self {
156        Self { source, deferred }
157    }
158
159    #[cfg(not(feature = "autodiff"))]
160    fn new(source: ShapeGuardError) -> Self {
161        Self { source }
162    }
163
164    /// Return the original typed shape-guard failure.
165    pub fn typed_source(&self) -> &ShapeGuardError {
166        &self.source
167    }
168
169    /// Consume this boundary error and return its original typed failure.
170    pub fn into_typed_source(self) -> ShapeGuardError {
171        self.source
172    }
173
174    #[cfg(feature = "autodiff")]
175    fn record_for_ad_boundary(&self) {
176        if let Ok(mut deferred) = self.deferred.lock() {
177            if deferred.is_none() {
178                *deferred = Some(self.source.clone());
179            }
180        }
181    }
182}
183
184impl std::fmt::Display for ShapeGuardFailure {
185    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
186        self.source.fmt(f)
187    }
188}
189
190impl std::error::Error for ShapeGuardFailure {
191    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
192        Some(&self.source)
193    }
194}
195
196impl PartialEq for ShapeGuardFailure {
197    fn eq(&self, other: &Self) -> bool {
198        self.source == other.source
199    }
200}
201
202impl Eq for ShapeGuardFailure {}
203
204/// Global metadata registry.
205///
206/// Stored as a tokenized stack per value key: duplicate scoped registrations
207/// shadow older metadata while they are live, and dropping scopes in any order
208/// removes only the matching token. `ShapeGuardContext::metadata_of` reaches into
209/// the registry lazily via [`lookup_global_metadata`] and caches the result into
210/// the context's local map.
211///
212/// Earlier designs either cloned the whole map up-front into each AD
213/// `ShapeGuardContext` or kept the map in an `Arc` and cloned on every write.
214/// Both variants were quadratic across the monotonically growing registry and
215/// dominated oracle_replay runtime.
216static GLOBAL_METADATA: OnceLock<Mutex<GlobalMetadataMap>> = OnceLock::new();
217static NEXT_GLOBAL_METADATA_TOKEN: AtomicU64 = AtomicU64::new(0);
218
219fn global_metadata_registry() -> &'static Mutex<GlobalMetadataMap> {
220    GLOBAL_METADATA.get_or_init(|| Mutex::new(HashMap::new()))
221}
222
223/// Lifetime token for graph-scoped global metadata.
224///
225/// Dropping the last frontend owner of a traced graph drops this scope and
226/// releases the metadata keys that were registered for that graph graph.
227#[doc(hidden)]
228#[derive(Debug)]
229pub struct GlobalMetadataScope {
230    registrations: Vec<ScopedGlobalMetadataRegistration>,
231}
232
233impl Drop for GlobalMetadataScope {
234    fn drop(&mut self) {
235        release_scoped_global_metadata(&self.registrations);
236    }
237}
238
239/// Per-value tensor metadata used by AD rules.
240///
241/// Shape information is stored as per-axis [`ShapeExtent`] values. Callers must
242/// explicitly choose whether they need an exact shape or only a known bound.
243///
244/// # Examples
245///
246/// ```
247/// use tenferro_ops::{SymDim, TensorMeta};
248/// use tenferro_tensor::DType;
249///
250/// let meta = TensorMeta::exact(DType::F64, vec![SymDim::from(2usize), SymDim::from(3usize)]);
251/// assert_eq!(meta.rank(), 2);
252/// ```
253#[derive(Clone, Debug, PartialEq, Eq)]
254pub struct TensorMeta {
255    /// Element dtype of the tensor value.
256    pub dtype: DType,
257    /// Per-axis shape guarantees.
258    pub extents: Vec<ShapeExtent<SymDim>>,
259    /// Canonical identity of an externally defined scalar, when the value carries one.
260    ///
261    /// A semantic program's identity must be reproducible across processes, while an
262    /// externally defined scalar's tag is a process-local `TypeId`, so a traced value
263    /// whose dtype is external declares the stable name here.
264    pub scalar_identity: Option<&'static str>,
265}
266
267impl TensorMeta {
268    /// Construct metadata whose every axis is exact.
269    ///
270    /// # Examples
271    ///
272    /// ```
273    /// use tenferro_ops::{SymDim, TensorMeta};
274    /// use tenferro_tensor::DType;
275    ///
276    /// let meta = TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]);
277    /// assert_eq!(meta.exact_shape(), Some(vec![SymDim::from(4usize)]));
278    /// ```
279    pub fn exact(dtype: DType, shape: Vec<SymDim>) -> Self {
280        let extents = shape.iter().cloned().map(ShapeExtent::exact).collect();
281        Self {
282            dtype,
283            extents,
284            scalar_identity: None,
285        }
286    }
287
288    /// Construct metadata from per-axis extents.
289    ///
290    /// # Examples
291    ///
292    /// ```
293    /// use tenferro_ops::{ShapeExtent, SymDim, TensorMeta};
294    /// use tenferro_tensor::DType;
295    ///
296    /// let meta = TensorMeta::with_extents(
297    ///     DType::F64,
298    ///     vec![ShapeExtent::upper_bound(SymDim::from(8usize))],
299    /// );
300    /// assert_eq!(meta.exact_shape(), None);
301    /// ```
302    pub fn with_extents(dtype: DType, extents: Vec<ShapeExtent<SymDim>>) -> Self {
303        Self {
304            dtype,
305            extents,
306            scalar_identity: None,
307        }
308    }
309
310    /// Declare the canonical identity of an externally defined scalar.
311    ///
312    /// # Examples
313    ///
314    /// ```rust
315    /// use tenferro_ops::{SymDim, TensorMeta};
316    /// use tenferro_tensor::DType;
317    ///
318    /// let dtype = DType::External(std::any::TypeId::of::<f64>());
319    /// let meta = TensorMeta::exact(dtype, vec![SymDim::from(2usize)])
320    ///     .with_scalar_identity("example.scalar.v1");
321    /// assert_eq!(meta.scalar_identity(), Some("example.scalar.v1"));
322    /// ```
323    #[must_use]
324    pub fn with_scalar_identity(mut self, identity: &'static str) -> Self {
325        self.scalar_identity = Some(identity);
326        self
327    }
328
329    /// Return the declared identity of an externally defined scalar, if any.
330    ///
331    /// # Examples
332    ///
333    /// ```rust
334    /// use tenferro_ops::{SymDim, TensorMeta};
335    /// use tenferro_tensor::DType;
336    ///
337    /// let dtype = DType::External(std::any::TypeId::of::<f64>());
338    /// assert_eq!(TensorMeta::exact(dtype, vec![SymDim::from(1usize)]).scalar_identity(), None);
339    /// assert_eq!(
340    ///     TensorMeta::exact(dtype, vec![SymDim::from(1usize)])
341    ///         .with_scalar_identity("example.scalar.v1")
342    ///         .scalar_identity(),
343    ///     Some("example.scalar.v1")
344    /// );
345    /// ```
346    #[must_use]
347    pub const fn scalar_identity(&self) -> Option<&'static str> {
348        self.scalar_identity
349    }
350
351    /// Return the tensor rank known by this metadata record.
352    pub fn rank(&self) -> usize {
353        self.extents.len()
354    }
355
356    /// Return the per-axis shape guarantees.
357    ///
358    /// # Examples
359    ///
360    /// ```
361    /// use tenferro_ops::{SymDim, TensorMeta};
362    /// use tenferro_tensor::DType;
363    ///
364    /// let meta = TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]);
365    /// assert_eq!(meta.extents().len(), 1);
366    /// ```
367    pub fn extents(&self) -> &[ShapeExtent<SymDim>] {
368        &self.extents
369    }
370
371    /// Return the shape only when every axis is exact.
372    ///
373    /// # Examples
374    ///
375    /// ```
376    /// use tenferro_ops::{ShapeExtent, SymDim, TensorMeta};
377    /// use tenferro_tensor::DType;
378    ///
379    /// let meta = TensorMeta::with_extents(
380    ///     DType::F64,
381    ///     vec![ShapeExtent::upper_bound(SymDim::from(8usize))],
382    /// );
383    /// assert_eq!(meta.exact_shape(), None);
384    /// ```
385    pub fn exact_shape(&self) -> Option<Vec<SymDim>> {
386        self.extents
387            .iter()
388            .map(|extent| extent.as_exact().cloned())
389            .collect()
390    }
391
392    /// Return one known bound per axis when every axis has a bound.
393    ///
394    /// This is intentionally separate from [`TensorMeta::exact_shape`]: a bound
395    /// is not proof of the runtime size.
396    pub fn bound_shape(&self) -> Option<Vec<SymDim>> {
397        self.extents
398            .iter()
399            .map(|extent| extent.bound_expr().cloned())
400            .collect()
401    }
402}
403
404/// A recorded dimension comparison made during AD graph construction.
405///
406/// # Examples
407///
408/// ```
409/// use std::cmp::Ordering;
410/// use tenferro_ops::ShapeGuard;
411///
412/// let guard = ShapeGuard {
413///     dim_a: 5,
414///     dim_b: 3,
415///     ordering: Ordering::Greater,
416/// };
417///
418/// assert_eq!(guard.ordering, Ordering::Greater);
419/// ```
420#[derive(Clone, Debug, PartialEq, Eq)]
421pub struct ShapeGuard {
422    /// First dimension value, such as `m`.
423    pub dim_a: usize,
424    /// Second dimension value, such as `n`.
425    pub dim_b: usize,
426    /// The observed ordering `dim_a.cmp(&dim_b)`.
427    pub ordering: Ordering,
428}
429
430/// AD context providing dimension resolution, guard recording, and value metadata.
431///
432/// # Examples
433///
434/// ```
435/// use tenferro_ops::ShapeGuardContext;
436///
437/// let ctx = ShapeGuardContext::default();
438/// assert!(ctx.guards().is_empty());
439/// ```
440#[derive(Clone, Debug, Default)]
441pub struct ShapeGuardContext {
442    guards: Vec<ShapeGuard>,
443    metadata: MetadataMap,
444    shape_sources: HashMap<u64, ValueKey<StdTensorOp>>,
445    use_global_registry: bool,
446    local_keys: Option<Vec<ValueKey<StdTensorOp>>>,
447    #[cfg(feature = "autodiff")]
448    deferred_shape_error: Arc<Mutex<Option<ShapeGuardError>>>,
449    #[cfg(feature = "autodiff")]
450    extension_ad_dispatcher: Option<Arc<dyn ExtensionAdDispatcher>>,
451    #[cfg(feature = "autodiff")]
452    active_value_keys: Option<std::sync::Arc<std::collections::HashSet<ValueKey<StdTensorOp>>>>,
453    #[cfg(feature = "autodiff")]
454    transpose_primal_outputs: Option<Vec<ValueKey<StdTensorOp>>>,
455    #[cfg(feature = "autodiff")]
456    transpose_primal_outputs_used: bool,
457}
458
459impl ShapeGuardContext {
460    /// Create a context backed by the global metadata registry.
461    ///
462    /// Instead of cloning the entire global registry up-front (which used
463    /// to be O(N) per AD pass and quadratic across oracle_replay), the
464    /// context keeps a flag and lazily fetches entries from the shared
465    /// [`lookup_global_metadata`] on first miss, caching into its local
466    /// `metadata` map for subsequent reads within the same pass.
467    ///
468    /// # Examples
469    ///
470    /// ```
471    /// let ctx = tenferro_ops::ShapeGuardContext::with_global_metadata();
472    /// assert!(ctx.guards().is_empty());
473    /// ```
474    pub fn with_global_metadata() -> Self {
475        Self {
476            use_global_registry: true,
477            ..Self::default()
478        }
479    }
480
481    #[doc(hidden)]
482    /// Keep global-registry lookup enabled after a pass boundary.
483    ///
484    /// This is intentionally a no-op for cached entries: global metadata is
485    /// already read lazily on cache misses, and clearing the local cache would
486    /// also discard metadata inserted directly into this context.
487    pub fn refresh_global_metadata(&mut self) {
488        self.use_global_registry = true;
489    }
490
491    #[doc(hidden)]
492    pub fn insert_shape_source(&mut self, tensor_id: u64, key: ValueKey<StdTensorOp>) {
493        self.shape_sources.entry(tensor_id).or_insert(key);
494    }
495
496    #[doc(hidden)]
497    pub fn shape_source(&self, tensor_id: u64) -> Option<&ValueKey<StdTensorOp>> {
498        self.shape_sources.get(&tensor_id)
499    }
500
501    #[doc(hidden)]
502    #[cfg(feature = "autodiff")]
503    pub fn with_extension_ad_dispatcher(
504        mut self,
505        dispatcher: Arc<dyn ExtensionAdDispatcher>,
506    ) -> Self {
507        self.extension_ad_dispatcher = Some(dispatcher);
508        self
509    }
510
511    #[doc(hidden)]
512    #[cfg(feature = "autodiff")]
513    pub(crate) fn extension_ad_dispatcher(&self) -> Option<Arc<dyn ExtensionAdDispatcher>> {
514        self.extension_ad_dispatcher.as_ref().map(Arc::clone)
515    }
516
517    #[cfg(feature = "autodiff")]
518    pub fn with_linearize_active_values(
519        mut self,
520        keys: std::sync::Arc<std::collections::HashSet<ValueKey<StdTensorOp>>>,
521    ) -> Self {
522        self.active_value_keys = Some(keys);
523        self
524    }
525
526    /// Whether a primal value lies on a path from the current linearize targets.
527    ///
528    /// When no active set was attached, every value is treated as active so
529    /// existing callers keep the conservative full JVP graphs.
530    #[cfg(feature = "autodiff")]
531    pub fn is_value_active_in_linearize(&self, key: &ValueKey<StdTensorOp>) -> bool {
532        self.active_value_keys
533            .as_ref()
534            .is_none_or(|set| set.contains(key))
535    }
536
537    /// Primal output keys for the operation currently being transposed.
538    ///
539    /// Primary-mode extension transpose rules such as `Eigh` use these to reuse
540    /// forward eigenvectors instead of recomputing a decomposition.
541    #[cfg(feature = "autodiff")]
542    pub fn set_transpose_primal_outputs(&mut self, keys: Option<Vec<ValueKey<StdTensorOp>>>) {
543        self.transpose_primal_outputs = keys;
544        self.transpose_primal_outputs_used = false;
545    }
546
547    /// Return the current primal outputs and mark them as consumed by this rule.
548    #[cfg(feature = "autodiff")]
549    pub fn transpose_primal_outputs(&mut self) -> Option<&[ValueKey<StdTensorOp>]> {
550        if self.transpose_primal_outputs.is_some() {
551            self.transpose_primal_outputs_used = true;
552        }
553        self.transpose_primal_outputs.as_deref()
554    }
555
556    #[cfg(feature = "autodiff")]
557    pub fn transpose_primal_outputs_were_used(&self) -> bool {
558        self.transpose_primal_outputs_used
559    }
560
561    /// Returns the guards recorded so far.
562    ///
563    /// # Examples
564    ///
565    /// ```
566    /// use tenferro_ops::ShapeGuardContext;
567    ///
568    /// let ctx = ShapeGuardContext::default();
569    /// assert_eq!(ctx.guards(), &[]);
570    /// ```
571    pub fn guards(&self) -> &[ShapeGuard] {
572        &self.guards
573    }
574
575    /// Clears all recorded guards.
576    ///
577    /// # Examples
578    ///
579    /// ```
580    /// use tenferro_ops::ShapeGuardContext;
581    ///
582    /// let mut ctx = ShapeGuardContext::default();
583    /// ctx.clear_guards();
584    /// assert!(ctx.guards().is_empty());
585    /// ```
586    pub fn clear_guards(&mut self) {
587        self.guards.clear();
588    }
589
590    /// Take the first typed shape-guard failure recorded while crossing an AD
591    /// callback boundary.
592    ///
593    /// AD callbacks expose only a message-bearing error. AD
594    /// frontends call this after the callback returns and attach the typed
595    /// value to their public runtime error.
596    #[doc(hidden)]
597    #[cfg(feature = "autodiff")]
598    pub fn take_deferred_shape_error(&mut self) -> Option<ShapeGuardError> {
599        self.deferred_shape_error
600            .lock()
601            .ok()
602            .and_then(|mut error| error.take())
603    }
604
605    /// Return the shape metadata for a value reference.
606    ///
607    /// # Examples
608    ///
609    /// ```
610    /// use computegraph::types::{ValueKey, ValueRef};
611    /// use tenferro_ops::input_key::TensorInputKey;
612    /// use tenferro_ops::std_tensor_op::StdTensorOp;
613    /// use tenferro_ops::{ShapeGuardContext, SymDim, TensorMeta};
614    /// use tenferro_tensor::DType;
615    ///
616    /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
617    /// let value = ValueRef::External(key.clone());
618    /// let mut ctx = ShapeGuardContext::default();
619    /// ctx.insert_metadata(key, TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]));
620    ///
621    /// let shape = ctx.shape_of(&value).unwrap();
622    /// assert_eq!(shape, &[SymDim::from(4usize)]);
623    /// ```
624    ///
625    /// # Errors
626    ///
627    /// Returns [`ShapeGuardError`] when the value cannot be resolved, metadata
628    /// is missing, or the metadata does not describe an exact shape.
629    pub fn shape_of(&mut self, val: &ValueRef<StdTensorOp>) -> ShapeGuardResult<Vec<SymDim>> {
630        let key = self.resolve_key(val)?.clone();
631        self.ensure_metadata_loaded(&key);
632        let meta = self.metadata.get(&key).ok_or_else(|| {
633            self.shape_guard_failure(ShapeGuardError::MissingMetadata { key: key.clone() })
634        })?;
635        meta.exact_shape()
636            .ok_or_else(|| self.shape_guard_failure(ShapeGuardError::NonExactShape { key }))
637    }
638
639    /// Return the rank for a value reference without requiring exact extents.
640    ///
641    /// Use this when an AD rule only needs axis count or needs to build
642    /// runtime-shape references. Calling [`ShapeGuardContext::shape_of`] in those
643    /// cases would reject valid values such as `DynamicTruncate` outputs whose
644    /// runtime extent is known only as an upper bound.
645    ///
646    /// # Examples
647    ///
648    /// ```
649    /// use computegraph::types::{ValueKey, ValueRef};
650    /// use tenferro_ops::input_key::TensorInputKey;
651    /// use tenferro_ops::std_tensor_op::StdTensorOp;
652    /// use tenferro_ops::{ShapeExtent, ShapeGuardContext, SymDim, TensorMeta};
653    /// use tenferro_tensor::DType;
654    ///
655    /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
656    /// let value = ValueRef::External(key.clone());
657    /// let mut ctx = ShapeGuardContext::default();
658    /// ctx.insert_metadata(
659    ///     key,
660    ///     TensorMeta::with_extents(DType::F64, vec![ShapeExtent::upper_bound(SymDim::from(8usize))]),
661    /// );
662    ///
663    /// assert_eq!(ctx.rank_of(&value).unwrap(), 1);
664    /// ```
665    ///
666    /// # Errors
667    ///
668    /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
669    /// metadata is unavailable.
670    pub fn rank_of(&mut self, val: &ValueRef<StdTensorOp>) -> ShapeGuardResult<usize> {
671        self.metadata_of(val).map(TensorMeta::rank)
672    }
673
674    /// Return per-axis shape guarantees for a value reference.
675    ///
676    /// # Examples
677    ///
678    /// ```
679    /// use computegraph::types::{ValueKey, ValueRef};
680    /// use tenferro_ops::input_key::TensorInputKey;
681    /// use tenferro_ops::std_tensor_op::StdTensorOp;
682    /// use tenferro_ops::{ShapeExtent, ShapeGuardContext, SymDim, TensorMeta};
683    /// use tenferro_tensor::DType;
684    ///
685    /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
686    /// let value = ValueRef::External(key.clone());
687    /// let mut ctx = ShapeGuardContext::default();
688    /// ctx.insert_metadata(
689    ///     key,
690    ///     TensorMeta::with_extents(DType::F64, vec![ShapeExtent::upper_bound(SymDim::from(8usize))]),
691    /// );
692    ///
693    /// let extents = ctx.extents_of(&value).unwrap();
694    /// assert_eq!(extents[0], ShapeExtent::upper_bound(SymDim::from(8usize)));
695    /// ```
696    ///
697    /// # Errors
698    ///
699    /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
700    /// metadata is unavailable.
701    pub fn extents_of(
702        &mut self,
703        val: &ValueRef<StdTensorOp>,
704    ) -> ShapeGuardResult<&[ShapeExtent<SymDim>]> {
705        self.metadata_of(val).map(TensorMeta::extents)
706    }
707
708    /// Return the exact shape for a value reference, if all axes are exact.
709    ///
710    /// # Examples
711    ///
712    /// ```
713    /// use computegraph::types::{ValueKey, ValueRef};
714    /// use tenferro_ops::input_key::TensorInputKey;
715    /// use tenferro_ops::std_tensor_op::StdTensorOp;
716    /// use tenferro_ops::{ShapeExtent, ShapeGuardContext, SymDim, TensorMeta};
717    /// use tenferro_tensor::DType;
718    ///
719    /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
720    /// let value = ValueRef::External(key.clone());
721    /// let mut ctx = ShapeGuardContext::default();
722    /// ctx.insert_metadata(
723    ///     key,
724    ///     TensorMeta::with_extents(DType::F64, vec![ShapeExtent::upper_bound(SymDim::from(8usize))]),
725    /// );
726    ///
727    /// let maybe_shape = ctx.exact_shape_of(&value).unwrap();
728    /// assert_eq!(maybe_shape, None);
729    /// ```
730    ///
731    /// # Errors
732    ///
733    /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
734    /// metadata is unavailable.
735    pub fn exact_shape_of(
736        &mut self,
737        val: &ValueRef<StdTensorOp>,
738    ) -> ShapeGuardResult<Option<Vec<SymDim>>> {
739        self.metadata_of(val).map(TensorMeta::exact_shape)
740    }
741
742    #[doc(hidden)]
743    pub fn shape_if_available(&mut self, val: &ValueRef<StdTensorOp>) -> Option<Vec<SymDim>> {
744        self.metadata_if_available(val)
745            .and_then(TensorMeta::exact_shape)
746    }
747
748    /// Return the dtype metadata for a value reference.
749    ///
750    /// # Examples
751    ///
752    /// ```
753    /// use computegraph::types::{ValueKey, ValueRef};
754    /// use tenferro_ops::input_key::TensorInputKey;
755    /// use tenferro_ops::std_tensor_op::StdTensorOp;
756    /// use tenferro_ops::{ShapeGuardContext, SymDim, TensorMeta};
757    /// use tenferro_tensor::DType;
758    ///
759    /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
760    /// let value = ValueRef::External(key.clone());
761    /// let mut ctx = ShapeGuardContext::default();
762    /// ctx.insert_metadata(key, TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]));
763    ///
764    /// let dtype = ctx.dtype_of(&value).unwrap();
765    /// assert_eq!(dtype, DType::F64);
766    /// ```
767    ///
768    /// # Errors
769    ///
770    /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
771    /// metadata is unavailable.
772    pub fn dtype_of(&mut self, val: &ValueRef<StdTensorOp>) -> ShapeGuardResult<DType> {
773        self.metadata_of(val).map(|meta| meta.dtype)
774    }
775
776    /// Return the complete metadata record for a value reference.
777    ///
778    /// # Examples
779    ///
780    /// ```
781    /// use computegraph::types::{ValueKey, ValueRef};
782    /// use tenferro_ops::input_key::TensorInputKey;
783    /// use tenferro_ops::std_tensor_op::StdTensorOp;
784    /// use tenferro_ops::{ShapeGuardContext, SymDim, TensorMeta};
785    /// use tenferro_tensor::DType;
786    ///
787    /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
788    /// let value = ValueRef::External(key.clone());
789    /// let mut ctx = ShapeGuardContext::default();
790    /// ctx.insert_metadata(key, TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]));
791    ///
792    /// let meta = ctx.metadata_of(&value).unwrap();
793    /// assert_eq!(meta.dtype, DType::F64);
794    /// ```
795    ///
796    /// # Errors
797    ///
798    /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
799    /// metadata is unavailable.
800    pub fn metadata_of(&mut self, val: &ValueRef<StdTensorOp>) -> ShapeGuardResult<&TensorMeta> {
801        let key = self.resolve_key(val)?.clone();
802        self.ensure_metadata_loaded(&key);
803        self.metadata
804            .get(&key)
805            .ok_or_else(|| self.shape_guard_failure(ShapeGuardError::MissingMetadata { key }))
806    }
807
808    #[doc(hidden)]
809    pub fn metadata_if_available(&mut self, val: &ValueRef<StdTensorOp>) -> Option<&TensorMeta> {
810        let key = self.resolve_key_if_available(val)?.clone();
811        self.ensure_metadata_loaded(&key);
812        self.metadata.get(&key)
813    }
814
815    #[doc(hidden)]
816    pub fn attach_graph(&mut self, graph: &Graph<StdTensorOp>) {
817        self.local_keys = Some(graph.values().iter().map(|node| node.key.clone()).collect());
818    }
819
820    #[doc(hidden)]
821    pub fn insert_metadata(&mut self, key: ValueKey<StdTensorOp>, meta: TensorMeta) {
822        self.metadata.insert(key, meta);
823    }
824
825    #[doc(hidden)]
826    pub fn extend_metadata<I>(&mut self, entries: I)
827    where
828        I: IntoIterator<Item = (ValueKey<StdTensorOp>, TensorMeta)>,
829    {
830        self.metadata.extend(entries);
831    }
832
833    fn resolve_key_if_available<'a>(
834        &'a self,
835        val: &'a ValueRef<StdTensorOp>,
836    ) -> Option<&'a ValueKey<StdTensorOp>> {
837        match val {
838            ValueRef::External(key) => Some(key),
839            ValueRef::Local(local_id) => self
840                .local_keys
841                .as_ref()
842                .and_then(|keys| keys.get(*local_id)),
843        }
844    }
845
846    fn resolve_key<'a>(
847        &'a self,
848        val: &'a ValueRef<StdTensorOp>,
849    ) -> ShapeGuardResult<&'a ValueKey<StdTensorOp>> {
850        match val {
851            ValueRef::External(key) => Ok(key),
852            ValueRef::Local(local_id) if self.local_keys.is_none() => Err(self
853                .shape_guard_failure(ShapeGuardError::LocalWithoutAttachedGraph {
854                    local_id: *local_id,
855                })),
856            ValueRef::Local(local_id) => self
857                .local_keys
858                .as_ref()
859                .and_then(|keys| keys.get(*local_id))
860                .ok_or_else(|| {
861                    self.shape_guard_failure(ShapeGuardError::LocalOutOfBounds {
862                        local_id: *local_id,
863                    })
864                }),
865        }
866    }
867
868    #[cfg(feature = "autodiff")]
869    fn shape_guard_failure(&self, source: ShapeGuardError) -> ShapeGuardFailure {
870        ShapeGuardFailure::new(source, Arc::clone(&self.deferred_shape_error))
871    }
872
873    #[cfg(not(feature = "autodiff"))]
874    fn shape_guard_failure(&self, source: ShapeGuardError) -> ShapeGuardFailure {
875        ShapeGuardFailure::new(source)
876    }
877
878    fn ensure_metadata_loaded(&mut self, key: &ValueKey<StdTensorOp>) {
879        if !self.metadata.contains_key(key) && self.use_global_registry {
880            if let Ok(Some(meta)) = lookup_global_metadata(key) {
881                self.metadata.insert(key.clone(), meta);
882            }
883        }
884    }
885}
886
887/// Look up a single metadata entry from the global registry.
888///
889/// Locks the registry briefly for a single `HashMap::get` + clone.
890///
891/// # Examples
892///
893/// ```
894/// use computegraph::types::ValueKey;
895/// use tenferro_ops::ad::context::lookup_global_metadata;
896/// use tenferro_ops::input_key::TensorInputKey;
897/// use tenferro_ops::std_tensor_op::StdTensorOp;
898///
899/// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 99 });
900/// let meta = lookup_global_metadata(&key).unwrap();
901/// assert!(meta.is_none());
902/// ```
903///
904/// # Errors
905///
906/// Returns [`MetadataRegistryError::LockPoisoned`] when the global metadata
907/// registry lock is poisoned.
908pub fn lookup_global_metadata(
909    key: &ValueKey<StdTensorOp>,
910) -> Result<Option<TensorMeta>, MetadataRegistryError> {
911    let guard = global_metadata_registry()
912        .lock()
913        .map_err(|_| MetadataRegistryError::LockPoisoned)?;
914    Ok(guard
915        .get(key)
916        .and_then(|entry| entry.stack.last())
917        .map(|registration| registration.meta.clone()))
918}
919
920#[doc(hidden)]
921///
922/// # Errors
923///
924/// Returns [`MetadataRegistryError::LockPoisoned`] when the global metadata
925/// registry lock is poisoned.
926pub fn register_scoped_global_metadata_batch<I>(
927    entries: I,
928) -> Result<GlobalMetadataScope, MetadataRegistryError>
929where
930    I: IntoIterator<Item = (ValueKey<StdTensorOp>, TensorMeta)>,
931{
932    let mut guard = global_metadata_registry()
933        .lock()
934        .map_err(|_| MetadataRegistryError::LockPoisoned)?;
935    let mut registrations = Vec::new();
936    for (key, meta) in entries {
937        let token = NEXT_GLOBAL_METADATA_TOKEN.fetch_add(1, AtomicOrdering::Relaxed);
938        let entry = guard
939            .entry(key.clone())
940            .or_insert_with(|| GlobalMetadataEntry { stack: Vec::new() });
941        entry.stack.push(GlobalMetadataRegistration { token, meta });
942        registrations.push(ScopedGlobalMetadataRegistration { key, token });
943    }
944    Ok(GlobalMetadataScope { registrations })
945}
946
947fn release_scoped_global_metadata(registrations: &[ScopedGlobalMetadataRegistration]) {
948    let Ok(mut guard) = global_metadata_registry().lock() else {
949        // Drop cannot return an error. Failing closed here avoids reading or
950        // mutating data from a poisoned registry at the cost of leaking entries
951        // until process exit.
952        return;
953    };
954    for registration in registrations {
955        let should_remove = if let Some(entry) = guard.get_mut(&registration.key) {
956            if let Some(position) = entry
957                .stack
958                .iter()
959                .rposition(|candidate| candidate.token == registration.token)
960            {
961                entry.stack.remove(position);
962            }
963            entry.stack.is_empty()
964        } else {
965            false
966        };
967        if should_remove {
968            guard.remove(&registration.key);
969        }
970    }
971}
972
973/// Resolve a [`DimExpr`] to a concrete `usize`.
974#[doc(hidden)]
975pub fn resolve_dim(dim: &DimExpr) -> Result<usize, DimExprEvalError> {
976    dim.eval(&[])
977}
978
979/// Resolve matrix dimensions and record their ordering as a guard.
980#[doc(hidden)]
981pub fn resolve_and_guard(
982    m: &DimExpr,
983    n: &DimExpr,
984    ctx: &mut ShapeGuardContext,
985) -> Result<(usize, usize), DimExprEvalError> {
986    let m_size = resolve_dim(m)?;
987    let n_size = resolve_dim(n)?;
988    ctx.guards.push(ShapeGuard {
989        dim_a: m_size,
990        dim_b: n_size,
991        ordering: m_size.cmp(&n_size),
992    });
993    Ok((m_size, n_size))
994}