tenferro_runtime/error.rs
1//! Error types for the tenferro runtime crate.
2//!
3//! # Examples
4//!
5//! ```rust
6//! use tenferro_runtime::error::{Error, ErrorPhase};
7//!
8//! let err = Error::invalid_argument(
9//! "einsum",
10//! ErrorPhase::GraphBuild,
11//! "subscripts",
12//! "bad label",
13//! );
14//! assert!(err.to_string().contains("bad label"));
15//! ```
16
17use std::error::Error as StdError;
18use std::sync::atomic::{AtomicUsize, Ordering};
19
20use tenferro_ops::{dim_expr::DimExprEvalError, ShapeRelation, SymDimConversionError};
21use tenferro_tensor::{DType, ErrorKind, ValidationError, ValidationKind};
22
23use crate::runtime::{PrepareError, UnsupportedReason};
24
25static NEXT_CONTEXT_ID: AtomicUsize = AtomicUsize::new(1);
26
27/// Boxed source used when a runtime registry or compiler subsystem crosses
28/// the runtime error boundary with a concrete error owned by another crate.
29pub type BoxError = Box<dyn StdError + Send + Sync + 'static>;
30
31/// Borrowed, stable classification of a runtime preparation or execution failure.
32///
33/// The view retains references to the original error payloads and performs no
34/// message formatting or allocation. Match a wildcard because this enum is
35/// non-exhaustive.
36///
37/// # Examples
38///
39/// ```rust
40/// use tenferro_runtime::{Error, ErrorPhase, RuntimeFailureReasonRef};
41///
42/// let error = Error::unsupported("compare", ErrorPhase::Compile, "not available");
43/// assert!(matches!(
44/// error.reason(),
45/// RuntimeFailureReasonRef::UnsupportedOperation { operation: "compare" }
46/// ));
47/// ```
48#[derive(Clone, Copy, Debug, PartialEq)]
49#[non_exhaustive]
50pub enum RuntimeFailureReasonRef<'a> {
51 /// An operation belongs to an extension family with no installed module.
52 MissingExtension { family: &'a str },
53 /// No prepared engine accepts an input at its physical placement.
54 NoInputIngress {
55 input_index: usize,
56 placement: &'a tenferro_tensor::Placement,
57 },
58 /// A provider or tensor backend does not implement an operation.
59 UnsupportedOperation { operation: &'a str },
60 /// A failure has no stable structured classification.
61 Other,
62}
63
64/// Phase at which a runtime failure was discovered.
65///
66/// The phase is independent from [`ErrorKind`]: the same validation fact can
67/// be discovered while building a graph, compiling it for concrete inputs,
68/// or executing a compiled program.
69///
70/// # Examples
71///
72/// ```rust
73/// use tenferro_runtime::ErrorPhase;
74///
75/// assert_ne!(ErrorPhase::GraphBuild, ErrorPhase::Execution);
76/// ```
77#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
78#[non_exhaustive]
79pub enum ErrorPhase {
80 /// A caller-controlled graph construction check failed.
81 GraphBuild,
82 /// Shape inference or lowering discovered the failure.
83 Compile,
84 /// Input binding or backend execution discovered the failure.
85 Execution,
86}
87
88/// Typed reason that a symbolic shape constraint could not be evaluated.
89///
90/// # Examples
91///
92/// ```rust
93/// use tenferro_runtime::ShapeConstraintEvalError;
94///
95/// let cause = ShapeConstraintEvalError::MissingInput {
96/// input_idx: 2,
97/// input_count: 1,
98/// };
99/// assert_eq!(
100/// cause.to_string(),
101/// "shape expression references input 2, but only 1 inputs were provided"
102/// );
103/// ```
104#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
105pub enum ShapeConstraintEvalError {
106 /// An expression referenced an input shape that was not supplied.
107 #[error(
108 "shape expression references input {input_idx}, but only {input_count} inputs were provided"
109 )]
110 MissingInput {
111 /// Referenced input index.
112 input_idx: usize,
113 /// Number of supplied input shapes.
114 input_count: usize,
115 },
116 /// An expression referenced an axis outside the selected input's rank.
117 #[error("shape expression references input {input_idx} axis {axis}, but its rank is {rank}")]
118 AxisOutOfBounds {
119 /// Referenced input index.
120 input_idx: usize,
121 /// Referenced axis.
122 axis: usize,
123 /// Rank of the selected input.
124 rank: usize,
125 },
126 /// Checked dimension arithmetic overflowed `usize`.
127 #[error("shape expression arithmetic overflowed")]
128 Overflow,
129 /// Checked dimension subtraction underflowed `usize`.
130 #[error("shape expression subtraction underflowed")]
131 Underflow,
132 /// A floor-division divisor evaluated to zero.
133 #[error("shape expression divided by zero")]
134 DivisionByZero,
135}
136
137impl From<DimExprEvalError> for ShapeConstraintEvalError {
138 fn from(error: DimExprEvalError) -> Self {
139 match error {
140 DimExprEvalError::InputOutOfBounds {
141 input_idx,
142 input_count,
143 } => Self::MissingInput {
144 input_idx,
145 input_count,
146 },
147 DimExprEvalError::AxisOutOfBounds {
148 input_idx,
149 axis,
150 rank,
151 } => Self::AxisOutOfBounds {
152 input_idx,
153 axis,
154 rank,
155 },
156 DimExprEvalError::AddOverflow { .. } | DimExprEvalError::MulOverflow { .. } => {
157 Self::Overflow
158 }
159 DimExprEvalError::SubUnderflow { .. } => Self::Underflow,
160 DimExprEvalError::FloorDivByZero { .. } => Self::DivisionByZero,
161 }
162 }
163}
164
165/// Errors produced by einsum, eval, and other tenferro operations.
166///
167/// # Examples
168///
169/// ```rust
170/// use tenferro_runtime::error::{Error, ErrorPhase};
171///
172/// let err = Error::invalid_argument(
173/// "einsum",
174/// ErrorPhase::GraphBuild,
175/// "subscripts",
176/// "rank mismatch",
177/// );
178/// ```
179#[derive(Debug, thiserror::Error)]
180pub enum Error {
181 /// A shared tensor validation fact, annotated with the runtime phase.
182 #[error("{op} ({phase:?}): {source}")]
183 Validation {
184 /// Public operation name.
185 op: &'static str,
186 /// Phase that discovered the validation fact.
187 phase: ErrorPhase,
188 /// Machine-readable validation payload.
189 #[source]
190 source: ValidationError,
191 },
192
193 /// A required input tensor is missing from the inputs map.
194 #[error("missing input: {0}")]
195 MissingInput(String),
196
197 /// Reverse-mode gradient requires a scalar output.
198 #[error("grad requires a scalar output, got shape {shape:?}")]
199 NonScalarGrad { shape: Vec<usize> },
200
201 /// The operation is known not to support the requested input or
202 /// configuration at the phase where it was requested.
203 #[error("{op} ({phase:?}) is unsupported: {message}")]
204 Unsupported {
205 /// Operation that does not provide the requested behavior.
206 op: &'static str,
207 /// Phase that established the unsupported combination.
208 phase: ErrorPhase,
209 /// Human-readable unsupported-operation detail.
210 message: String,
211 },
212
213 /// Runtime tensor execution failed in the backend layer.
214 #[error(transparent)]
215 TensorRuntime(#[from] tenferro_tensor::Error),
216
217 /// A backend refused to open an execution session; no operation ran.
218 #[error("backend session entry failed: {0}")]
219 SessionEntry(#[source] tenferro_tensor::SessionEntryError),
220
221 /// A typed extension-domain error crossed a runtime registry boundary.
222 #[error("extension {family} ({phase:?}) failed for {op}: {source}")]
223 Extension {
224 /// Operation that discovered the extension failure.
225 op: &'static str,
226 /// Phase that discovered the extension failure.
227 phase: ErrorPhase,
228 /// Stable extension family identifier.
229 family: &'static str,
230 /// Coarse classification supplied by the extension owner.
231 kind: ErrorKind,
232 /// Original extension-domain source.
233 #[source]
234 source: BoxError,
235 },
236
237 /// Executor, cache, registry, or device state is unavailable or invalid.
238 #[error("{op} ({phase:?}): runtime state failure: {message}")]
239 RuntimeState {
240 /// Operation whose state was unavailable.
241 op: &'static str,
242 /// Phase that discovered the invalid state.
243 phase: ErrorPhase,
244 /// Human-readable state detail.
245 message: String,
246 },
247
248 /// A runtime-state failure retaining a typed source.
249 #[error("{op} ({phase:?}): runtime state failure: {source}")]
250 RuntimeStateSource {
251 /// Operation whose state was unavailable.
252 op: &'static str,
253 /// Phase that discovered the invalid state.
254 phase: ErrorPhase,
255 /// Typed state source.
256 #[source]
257 source: BoxError,
258 },
259
260 /// A primary error with a second typed error retained as suppressed
261 /// metadata.
262 ///
263 /// The standard error source chain follows `primary`. The suppressed
264 /// error is intentionally exposed through [`Error::suppressed`] because
265 /// [`StdError::source`](std::error::Error::source) can represent only one
266 /// source without losing the primary error's semantics.
267 #[error("primary error: {primary}; suppressed error: {suppressed}")]
268 WithSuppressed {
269 /// The operation's primary failure and the standard error-chain source.
270 #[source]
271 primary: Box<Error>,
272 /// A typed secondary failure retained for diagnostics and recovery.
273 suppressed: Box<Error>,
274 },
275
276 /// A runtime event-domain provenance or admission contract failed.
277 #[error("event-domain operation failed: {source}")]
278 EventDomain {
279 /// Structured event-domain failure with expected/actual provenance.
280 #[from]
281 #[source]
282 source: crate::runtime::EventDomainError,
283 },
284
285 /// A `TracedTensor` supplied as a compiled-graph input binding is not a
286 /// placeholder (has attached data).
287 #[error(
288 "binding #{binding_index} is not a placeholder; \
289 only tensors built via input_concrete_shape / input_symbolic_shape \
290 can be bound"
291 )]
292 UnexpectedBinding { binding_index: usize },
293
294 /// A placeholder appearing in the graph has no binding supplied.
295 #[error("placeholder {input_key} has no runtime input binding")]
296 UnboundPlaceholder { input_key: String },
297
298 /// The number of ordered tensors supplied to a compiled graph is invalid.
299 #[error("compiled graph expects {expected} ordered inputs, got {actual}")]
300 GraphInputCountMismatch { expected: usize, actual: usize },
301
302 /// The same placeholder was bound more than once in the `bindings` slice.
303 #[error("placeholder {input_key} was bound more than once")]
304 DuplicateBinding { input_key: String },
305
306 /// A binding tensor's dtype does not match the placeholder's dtype.
307 #[error("binding dtype mismatch for placeholder: expected {expected:?}, got {actual:?}")]
308 PlaceholderDtypeMismatch { expected: DType, actual: DType },
309
310 /// A binding tensor's shape does not match an `input_concrete_shape`
311 /// placeholder's fixed shape.
312 #[error(
313 "binding shape mismatch for concrete-shape placeholder: \
314 expected {expected:?}, got {actual:?}"
315 )]
316 PlaceholderShapeMismatch {
317 expected: Vec<usize>,
318 actual: Vec<usize>,
319 },
320
321 /// A binding tensor dimension exceeds a semantic input's declared bound.
322 #[error("binding dimension {axis} exceeds semantic input upper bound {bound}: got {actual}")]
323 PlaceholderShapeBoundExceeded {
324 /// Axis whose runtime extent exceeded the bound.
325 axis: usize,
326 /// Evaluated upper bound.
327 bound: usize,
328 /// Runtime extent.
329 actual: usize,
330 },
331
332 /// A binding tensor's rank does not match an `input_symbolic_shape`
333 /// placeholder's declared rank.
334 #[error(
335 "binding rank mismatch for symbolic-shape placeholder: \
336 expected rank {expected}, got rank {actual}"
337 )]
338 PlaceholderRankMismatch { expected: usize, actual: usize },
339
340 /// Operation attempted to mix tensors from different eager contexts.
341 #[error(
342 "tensors belong to different eager AD contexts ({lhs} vs {rhs}); \
343 detach into the target context before combining them"
344 )]
345 ContextMismatch { lhs: ContextId, rhs: ContextId },
346
347 /// An AD transform requires a primitive or extension rule that is not
348 /// registered for the requested operation.
349 #[error("unsupported {transform} AD rule for {op}")]
350 UnsupportedAdRule {
351 /// AD transform that requested the rule, such as `grad` or `backward`.
352 transform: &'static str,
353 /// Operation or extension family identifier that has no applicable rule.
354 op: String,
355 },
356
357 /// A typed AD rule source that crossed an external message-only callback.
358 #[error("{transform} AD rule failed: {source}")]
359 AdRuleSource {
360 /// AD transform that requested the rule.
361 transform: &'static str,
362 /// Original typed source from the AD rule context.
363 #[source]
364 source: BoxError,
365 },
366
367 /// A symbolic extension shape equality evaluated to unequal dimensions.
368 #[error(
369 "extension family {family:?} shape constraint at instruction {instruction_index:?} failed: {lhs_expr} ({lhs_value}) {relation:?} {rhs_expr} ({rhs_value})"
370 )]
371 ShapeConstraintViolation {
372 /// Stable extension family identifier.
373 family: &'static str,
374 /// Stable compiled instruction provenance, when assigned.
375 instruction_index: Option<usize>,
376 /// Shape relation that failed.
377 relation: ShapeRelation,
378 /// Normalized left-hand expression.
379 lhs_expr: String,
380 /// Normalized right-hand expression.
381 rhs_expr: String,
382 /// Concrete left-hand value.
383 lhs_value: usize,
384 /// Concrete right-hand value.
385 rhs_value: usize,
386 },
387
388 /// A symbolic extension shape expression could not be evaluated safely.
389 #[error(
390 "extension family {family:?} shape constraint at instruction {instruction_index:?} could not evaluate {expression} for {relation:?}: {cause}"
391 )]
392 ShapeConstraintEvaluation {
393 /// Stable extension family identifier.
394 family: &'static str,
395 /// Stable compiled instruction provenance, when assigned.
396 instruction_index: Option<usize>,
397 /// Shape relation whose expression failed.
398 relation: ShapeRelation,
399 /// Normalized expression that failed.
400 expression: String,
401 /// Typed evaluation failure.
402 #[source]
403 cause: ShapeConstraintEvalError,
404 },
405
406 /// A symbolic dimension could not be converted into the graph's local
407 /// dimension-expression vocabulary.
408 #[error("{op} ({phase:?}): symbolic shape conversion failed: {source}")]
409 SymbolicShapeConversion {
410 /// Operation that requested the symbolic shape conversion.
411 op: &'static str,
412 /// Phase that discovered the invalid symbolic reference.
413 phase: ErrorPhase,
414 /// Typed symbolic-dimension conversion failure.
415 #[source]
416 source: SymDimConversionError,
417 },
418
419 /// A runtime dimension expression could not be evaluated for concrete
420 /// input shapes.
421 #[error("runtime shape expression {expression} could not evaluate: {cause}")]
422 ShapeExpressionEvaluation {
423 /// Expression that failed during execution.
424 expression: String,
425 /// Typed evaluation failure.
426 #[source]
427 cause: ShapeConstraintEvalError,
428 },
429
430 /// An unexpected internal error.
431 #[error("internal error: {0}")]
432 Internal(String),
433}
434
435impl Error {
436 /// Wrap a shared validation payload with its operation and discovery
437 /// phase.
438 ///
439 /// # Errors
440 ///
441 /// This constructor does not fail; callers receive the returned
442 /// [`Error`] value and can inspect its [`Error::kind`] and
443 /// [`Error::phase`].
444 ///
445 /// # Examples
446 ///
447 /// ```rust
448 /// use tenferro_runtime::{Error, ErrorPhase};
449 /// use tenferro_tensor::{ErrorKind, ShapeMismatch, ValidationKind};
450 ///
451 /// let error = Error::validation(
452 /// "reshape",
453 /// ErrorPhase::GraphBuild,
454 /// ShapeMismatch::ReshapeElementCount { from: 2, to: 3 }.into(),
455 /// );
456 /// assert_eq!(
457 /// error.kind(),
458 /// ErrorKind::Validation(ValidationKind::ShapeMismatch)
459 /// );
460 /// assert_eq!(error.phase(), Some(ErrorPhase::GraphBuild));
461 /// ```
462 pub fn validation(op: &'static str, phase: ErrorPhase, source: ValidationError) -> Self {
463 Self::Validation { op, phase, source }
464 }
465
466 /// Construct a validation error for a caller-controlled argument whose
467 /// failure does not have a more specific shared payload.
468 ///
469 /// # Examples
470 ///
471 /// ```rust
472 /// use tenferro_runtime::{Error, ErrorPhase};
473 ///
474 /// let error = Error::invalid_argument(
475 /// "broadcast_in_dim",
476 /// ErrorPhase::GraphBuild,
477 /// "dims",
478 /// "dimension mapping has the wrong length",
479 /// );
480 /// assert!(matches!(error, Error::Validation { .. }));
481 /// ```
482 pub fn invalid_argument(
483 op: &'static str,
484 phase: ErrorPhase,
485 argument: &'static str,
486 message: impl Into<String>,
487 ) -> Self {
488 Self::validation(
489 op,
490 phase,
491 ValidationError::InvalidArgument {
492 argument,
493 message: message.into(),
494 },
495 )
496 }
497
498 /// Construct a dtype-mismatch validation error using the runtime dtype
499 /// vocabulary.
500 ///
501 /// # Examples
502 ///
503 /// ```rust
504 /// use tenferro_runtime::{DType, Error, ErrorPhase};
505 /// use tenferro_tensor::{ErrorKind, ValidationKind};
506 ///
507 /// let error = Error::dtype_mismatch(
508 /// "add",
509 /// ErrorPhase::GraphBuild,
510 /// DType::F32,
511 /// DType::F64,
512 /// );
513 /// assert_eq!(error.kind(), ErrorKind::Validation(ValidationKind::DTypeMismatch));
514 /// ```
515 pub fn dtype_mismatch(
516 op: &'static str,
517 phase: ErrorPhase,
518 expected: DType,
519 actual: DType,
520 ) -> Self {
521 Self::validation(
522 op,
523 phase,
524 ValidationError::DTypeMismatch { expected, actual },
525 )
526 }
527
528 /// Preserve a typed extension-domain source at the runtime boundary.
529 ///
530 /// # Examples
531 ///
532 /// ```rust
533 /// use std::error::Error as _;
534 /// use tenferro_runtime::{Error, ErrorPhase};
535 /// use tenferro_tensor::ErrorKind;
536 ///
537 /// let source = std::io::Error::new(std::io::ErrorKind::Other, "extension failed");
538 /// let error = Error::extension(
539 /// "einsum",
540 /// ErrorPhase::GraphBuild,
541 /// "example.extension.v1",
542 /// ErrorKind::RuntimeState,
543 /// source,
544 /// );
545 /// assert_eq!(error.kind(), ErrorKind::RuntimeState);
546 /// assert!(error.source().is_some());
547 /// ```
548 pub fn extension<E>(
549 op: &'static str,
550 phase: ErrorPhase,
551 family: &'static str,
552 kind: ErrorKind,
553 source: E,
554 ) -> Self
555 where
556 E: StdError + Send + Sync + 'static,
557 {
558 Self::Extension {
559 op,
560 phase,
561 family,
562 kind,
563 source: Box::new(source),
564 }
565 }
566
567 /// Construct a runtime-state failure for an unavailable or invalid
568 /// executor, cache, registry, or device state.
569 ///
570 /// # Examples
571 ///
572 /// ```rust
573 /// use tenferro_runtime::{Error, ErrorPhase};
574 /// use tenferro_tensor::ErrorKind;
575 ///
576 /// let error = Error::runtime_state(
577 /// "executor",
578 /// ErrorPhase::Execution,
579 /// "the executor is not initialized",
580 /// );
581 /// assert_eq!(error.kind(), ErrorKind::RuntimeState);
582 /// ```
583 pub fn runtime_state(op: &'static str, phase: ErrorPhase, message: impl Into<String>) -> Self {
584 Self::RuntimeState {
585 op,
586 phase,
587 message: message.into(),
588 }
589 }
590
591 /// Preserve a typed source for an unavailable or invalid runtime state.
592 ///
593 /// # Examples
594 ///
595 /// ```rust
596 /// use std::error::Error as _;
597 /// use tenferro_runtime::{Error, ErrorPhase};
598 /// use tenferro_tensor::ErrorKind;
599 ///
600 /// let error = Error::runtime_state_source(
601 /// "metadata",
602 /// ErrorPhase::Compile,
603 /// std::io::Error::other("registry lock poisoned"),
604 /// );
605 /// assert_eq!(error.kind(), ErrorKind::RuntimeState);
606 /// assert!(error.source().is_some());
607 /// ```
608 pub fn runtime_state_source<E>(op: &'static str, phase: ErrorPhase, source: E) -> Self
609 where
610 E: StdError + Send + Sync + 'static,
611 {
612 Self::RuntimeStateSource {
613 op,
614 phase,
615 source: Box::new(source),
616 }
617 }
618
619 /// Retain a typed secondary error while preserving the primary error's
620 /// classification and standard source chain.
621 ///
622 /// # Examples
623 ///
624 /// ```rust
625 /// use tenferro_runtime::{Error, ErrorPhase};
626 ///
627 /// let error = Error::with_suppressed(
628 /// Error::unsupported("backend", ErrorPhase::Execution, "primary"),
629 /// Error::runtime_state("runtime", ErrorPhase::Execution, "suppressed"),
630 /// );
631 /// assert!(error.primary().is_some());
632 /// assert!(error.suppressed().is_some());
633 /// ```
634 pub fn with_suppressed(primary: Self, suppressed: Self) -> Self {
635 Self::WithSuppressed {
636 primary: Box::new(primary),
637 suppressed: Box::new(suppressed),
638 }
639 }
640
641 /// Return the primary error when this value is a suppressed-error
642 /// aggregate.
643 ///
644 /// # Examples
645 ///
646 /// ```rust
647 /// use tenferro_runtime::{Error, ErrorPhase};
648 ///
649 /// let error = Error::with_suppressed(
650 /// Error::unsupported("backend", ErrorPhase::Execution, "primary"),
651 /// Error::runtime_state("runtime", ErrorPhase::Execution, "suppressed"),
652 /// );
653 /// assert_eq!(error.primary().unwrap().phase(), Some(ErrorPhase::Execution));
654 /// ```
655 pub fn primary(&self) -> Option<&Self> {
656 match self {
657 Self::WithSuppressed { primary, .. } => Some(primary),
658 _ => None,
659 }
660 }
661
662 /// Return the typed suppressed error when this value is an aggregate.
663 ///
664 /// # Examples
665 ///
666 /// ```rust
667 /// use tenferro_runtime::{Error, ErrorPhase};
668 ///
669 /// let error = Error::with_suppressed(
670 /// Error::unsupported("backend", ErrorPhase::Execution, "primary"),
671 /// Error::runtime_state("runtime", ErrorPhase::Execution, "suppressed"),
672 /// );
673 /// assert_eq!(error.suppressed().unwrap().phase(), Some(ErrorPhase::Execution));
674 /// ```
675 pub fn suppressed(&self) -> Option<&Self> {
676 match self {
677 Self::WithSuppressed { suppressed, .. } => Some(suppressed),
678 _ => None,
679 }
680 }
681
682 /// Preserve a typed source returned by an AD rule through a callback
683 /// protocol that can carry only a rendered message.
684 ///
685 /// # Examples
686 ///
687 /// ```rust
688 /// use std::error::Error as _;
689 /// use tenferro_runtime::Error;
690 ///
691 /// let error = Error::ad_rule_source(
692 /// "jvp",
693 /// std::io::Error::other("shape metadata missing"),
694 /// );
695 /// assert!(error.source().is_some());
696 /// ```
697 pub fn ad_rule_source<E>(transform: &'static str, source: E) -> Self
698 where
699 E: StdError + Send + Sync + 'static,
700 {
701 Self::AdRuleSource {
702 transform,
703 source: Box::new(source),
704 }
705 }
706
707 /// Construct an operation-level unsupported error with an explicit
708 /// discovery phase.
709 ///
710 /// # Examples
711 ///
712 /// ```rust
713 /// use tenferro_runtime::{Error, ErrorPhase};
714 /// use tenferro_tensor::ErrorKind;
715 ///
716 /// let error = Error::unsupported(
717 /// "compare",
718 /// ErrorPhase::Compile,
719 /// "complex values have no total order",
720 /// );
721 /// assert_eq!(error.kind(), ErrorKind::Unsupported);
722 /// assert_eq!(error.phase(), Some(ErrorPhase::Compile));
723 /// ```
724 pub fn unsupported(op: &'static str, phase: ErrorPhase, message: impl Into<String>) -> Self {
725 Self::Unsupported {
726 op,
727 phase,
728 message: message.into(),
729 }
730 }
731
732 /// Return a borrowed stable reason for this runtime failure.
733 ///
734 /// Classification checks the primary error of [`Error::WithSuppressed`],
735 /// then walks the standard source chain. It never parses display text and
736 /// does not allocate; the original error, kind, phase, display, and source
737 /// chain remain unchanged.
738 ///
739 /// # Examples
740 ///
741 /// ```rust
742 /// use tenferro_runtime::{Error, ErrorPhase, RuntimeFailureReasonRef};
743 ///
744 /// let error = Error::unsupported("compare", ErrorPhase::Compile, "not available");
745 /// assert_eq!(
746 /// error.reason(),
747 /// RuntimeFailureReasonRef::UnsupportedOperation { operation: "compare" }
748 /// );
749 /// assert_eq!(Error::Internal("validation".into()).reason(), RuntimeFailureReasonRef::Other);
750 /// ```
751 pub fn reason(&self) -> RuntimeFailureReasonRef<'_> {
752 if let Self::WithSuppressed { primary, .. } = self {
753 return primary.reason();
754 }
755 if let Some(reason) = direct_reason(self) {
756 return reason;
757 }
758
759 let mut source = StdError::source(self);
760 while let Some(error) = source {
761 if let Some(reason) = error.downcast_ref::<Error>().and_then(direct_reason) {
762 return reason;
763 }
764 if let Some(reason) = prepare_reason(error) {
765 return reason;
766 }
767 source = StdError::source(error);
768 }
769 RuntimeFailureReasonRef::Other
770 }
771
772 /// Return the stable coarse classification of this runtime failure.
773 ///
774 /// # Examples
775 ///
776 /// ```rust
777 /// use tenferro_runtime::{Error, ErrorPhase};
778 /// use tenferro_tensor::{ErrorKind, ValidationError, ValidationKind};
779 ///
780 /// let error = Error::validation(
781 /// "transpose",
782 /// ErrorPhase::GraphBuild,
783 /// ValidationError::AxisOutOfBounds { axis: 2, rank: 2 },
784 /// );
785 /// assert_eq!(error.kind(), ErrorKind::Validation(ValidationKind::AxisOutOfBounds));
786 /// ```
787 pub fn kind(&self) -> ErrorKind {
788 match self {
789 Self::Validation { source, .. } => ErrorKind::Validation(source.kind()),
790 Self::MissingInput(_)
791 | Self::UnexpectedBinding { .. }
792 | Self::UnboundPlaceholder { .. }
793 | Self::DuplicateBinding { .. }
794 | Self::ContextMismatch { .. } => ErrorKind::RuntimeState,
795 Self::NonScalarGrad { .. } => ErrorKind::Validation(ValidationKind::InvalidArgument),
796 Self::GraphInputCountMismatch { .. } => {
797 ErrorKind::Validation(ValidationKind::InvalidArgument)
798 }
799 Self::Unsupported { .. } | Self::UnsupportedAdRule { .. } => ErrorKind::Unsupported,
800 Self::AdRuleSource { .. } => ErrorKind::Validation(ValidationKind::InvalidArgument),
801 Self::TensorRuntime(error) => error.kind(),
802 Self::SessionEntry(error) => error.kind(),
803 Self::Extension { kind, .. } => *kind,
804 Self::RuntimeState { .. }
805 | Self::RuntimeStateSource { .. }
806 | Self::EventDomain { .. } => ErrorKind::RuntimeState,
807 Self::WithSuppressed { primary, .. } => primary.kind(),
808 Self::PlaceholderDtypeMismatch { .. } => {
809 ErrorKind::Validation(ValidationKind::DTypeMismatch)
810 }
811 Self::PlaceholderShapeMismatch { .. } | Self::PlaceholderShapeBoundExceeded { .. } => {
812 ErrorKind::Validation(ValidationKind::ShapeMismatch)
813 }
814 Self::PlaceholderRankMismatch { .. } => {
815 ErrorKind::Validation(ValidationKind::RankMismatch)
816 }
817 Self::ShapeConstraintViolation { .. } => {
818 ErrorKind::Validation(ValidationKind::ShapeMismatch)
819 }
820 Self::ShapeConstraintEvaluation { .. } => {
821 ErrorKind::Validation(ValidationKind::InvalidArgument)
822 }
823 Self::SymbolicShapeConversion { .. } => {
824 ErrorKind::Validation(ValidationKind::InvalidArgument)
825 }
826 Self::ShapeExpressionEvaluation { .. } => {
827 ErrorKind::Validation(ValidationKind::InvalidArgument)
828 }
829 Self::Internal(_) => ErrorKind::Internal,
830 }
831 }
832
833 /// Return the discovery phase when this error has one.
834 ///
835 /// # Examples
836 ///
837 /// ```rust
838 /// use tenferro_runtime::{Error, ErrorPhase};
839 /// use tenferro_tensor::ValidationError;
840 ///
841 /// let error = Error::validation(
842 /// "reshape",
843 /// ErrorPhase::Compile,
844 /// ValidationError::RankMismatch { expected: 2, actual: 1 },
845 /// );
846 /// assert_eq!(error.phase(), Some(ErrorPhase::Compile));
847 /// ```
848 pub fn phase(&self) -> Option<ErrorPhase> {
849 match self {
850 Self::Validation { phase, .. } => Some(*phase),
851 Self::TensorRuntime(_) | Self::SessionEntry(_) => Some(ErrorPhase::Execution),
852 Self::Unsupported { phase, .. } => Some(*phase),
853 Self::Extension { phase, .. } => Some(*phase),
854 Self::RuntimeState { phase, .. } | Self::RuntimeStateSource { phase, .. } => {
855 Some(*phase)
856 }
857 Self::WithSuppressed { primary, .. } => primary.phase(),
858 Self::AdRuleSource { .. } => Some(ErrorPhase::GraphBuild),
859 Self::PlaceholderDtypeMismatch { .. }
860 | Self::PlaceholderShapeMismatch { .. }
861 | Self::PlaceholderRankMismatch { .. }
862 | Self::GraphInputCountMismatch { .. }
863 | Self::UnexpectedBinding { .. }
864 | Self::UnboundPlaceholder { .. }
865 | Self::DuplicateBinding { .. } => Some(ErrorPhase::Execution),
866 Self::EventDomain { .. } => Some(ErrorPhase::Execution),
867 Self::SymbolicShapeConversion { phase, .. } => Some(*phase),
868 Self::ShapeExpressionEvaluation { .. } => Some(ErrorPhase::Execution),
869 _ => None,
870 }
871 }
872}
873
874impl From<tenferro_tensor::SessionEntryError> for Error {
875 fn from(source: tenferro_tensor::SessionEntryError) -> Self {
876 Self::SessionEntry(source)
877 }
878}
879
880fn direct_reason(error: &Error) -> Option<RuntimeFailureReasonRef<'_>> {
881 match error {
882 Error::Unsupported { op, .. } => {
883 Some(RuntimeFailureReasonRef::UnsupportedOperation { operation: op })
884 }
885 Error::UnsupportedAdRule { op, .. } => {
886 Some(RuntimeFailureReasonRef::UnsupportedOperation { operation: op })
887 }
888 Error::TensorRuntime(error) => tensor_reason(error),
889 Error::Extension { family, kind, .. } if *kind == ErrorKind::Unsupported => {
890 Some(RuntimeFailureReasonRef::UnsupportedOperation { operation: family })
891 }
892 _ => None,
893 }
894}
895
896fn tensor_reason(error: &tenferro_tensor::Error) -> Option<RuntimeFailureReasonRef<'_>> {
897 match error {
898 tenferro_tensor::Error::Unsupported { op, .. }
899 | tenferro_tensor::Error::UnsupportedDType { op, .. }
900 | tenferro_tensor::Error::UnsupportedDTypeConversion { op, .. } => {
901 Some(RuntimeFailureReasonRef::UnsupportedOperation { operation: op })
902 }
903 _ => None,
904 }
905}
906
907fn prepare_reason<'a>(error: &'a (dyn StdError + 'static)) -> Option<RuntimeFailureReasonRef<'a>> {
908 let prepare = error.downcast_ref::<PrepareError>()?;
909 Some(match prepare {
910 PrepareError::MissingExtension { family_id } => {
911 RuntimeFailureReasonRef::MissingExtension { family: family_id }
912 }
913 PrepareError::NoInputIngress {
914 input_index,
915 placement,
916 } => RuntimeFailureReasonRef::NoInputIngress {
917 input_index: *input_index,
918 placement,
919 },
920 PrepareError::Unsupported {
921 reason: UnsupportedReason::Operation { operation },
922 } => RuntimeFailureReasonRef::UnsupportedOperation { operation },
923 _ => return None,
924 })
925}
926
927/// Opaque identifier for an eager AD runtime, used in [`Error::ContextMismatch`].
928#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
929pub struct ContextId(usize);
930
931impl ContextId {
932 /// Generate a fresh opaque runtime context identifier.
933 ///
934 /// Runtime implementations use this when constructing a new execution
935 /// context. The value is intentionally opaque and is only useful in error
936 /// reporting and equality checks.
937 pub fn fresh() -> Self {
938 let id = NEXT_CONTEXT_ID.fetch_add(1, Ordering::Relaxed);
939 Self(id)
940 }
941}
942
943impl std::fmt::Display for ContextId {
944 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
945 write!(f, "ctx@{:x}", self.0)
946 }
947}
948
949/// Result type alias for tenferro operations.
950pub type Result<T> = std::result::Result<T, Error>;
951
952#[cfg(test)]
953mod tests;