Skip to main content

tenferro_tensor/
error.rs

1//! Runtime error types for tensor execution.
2//!
3//! # Examples
4//!
5//! ```rust
6//! let error = tenferro_tensor::Error::shape_mismatch("add", [2], [3]);
7//! assert!(matches!(
8//!     error,
9//!     tenferro_tensor::Error::Validation { op: "add", .. }
10//! ));
11//! ```
12
13use std::error::Error as StdError;
14
15use tenferro_tensor_core::{ErrorKind, ValidationError};
16
17/// Boxed source used for backend and extension failures whose concrete type is
18/// owned by another crate or a vendor API.
19pub type BoxError = Box<dyn StdError + Send + Sync + 'static>;
20
21/// Runtime failures produced by tensor execution backends and helpers.
22///
23/// Validation failures retain the shared tensor-core payload as a typed source.
24/// Backend and extension failures retain opaque typed sources when one exists;
25/// text-only vendor failures use [`Error::BackendFailure`].
26///
27/// # Examples
28///
29/// ```rust
30/// let error = tenferro_tensor::Error::rank_mismatch("reshape", 2, 1);
31/// assert!(matches!(
32///     error,
33///     tenferro_tensor::Error::Validation { op: "reshape", .. }
34/// ));
35/// ```
36#[derive(Debug, thiserror::Error)]
37#[non_exhaustive]
38pub enum Error {
39    #[error("{op}: {source}")]
40    Validation {
41        op: &'static str,
42        #[source]
43        source: ValidationError,
44    },
45    #[error("{op}: unsupported dtype conversion from {from:?} to {to:?}: {message}")]
46    UnsupportedDTypeConversion {
47        op: &'static str,
48        from: crate::DType,
49        to: crate::DType,
50        message: String,
51    },
52    #[error("{op}: unsupported dtype {dtype:?}: {message}")]
53    UnsupportedDType {
54        op: &'static str,
55        dtype: crate::DType,
56        message: String,
57    },
58    #[error("{op}: unsupported operation: {message}")]
59    Unsupported { op: &'static str, message: String },
60    #[error("{op}: backend failure: {message}")]
61    BackendFailure { op: &'static str, message: String },
62    #[error("{op}: backend failure: {source}")]
63    BackendSource {
64        op: &'static str,
65        #[source]
66        source: BoxError,
67    },
68    #[error("{op}: I/O failure: {source}")]
69    IoSource {
70        op: &'static str,
71        #[source]
72        source: BoxError,
73    },
74    #[error("{op}: runtime state failure: {message}")]
75    RuntimeState { op: &'static str, message: String },
76    #[error("{op}: runtime state failure: {source}")]
77    RuntimeStateSource {
78        op: &'static str,
79        #[source]
80        source: BoxError,
81    },
82    #[error("{op}: host access failed: {source}")]
83    HostAccess {
84        op: &'static str,
85        #[source]
86        source: crate::HostAccessError,
87    },
88    #[error("{op}: extension {family} failed: {source}")]
89    Extension {
90        op: &'static str,
91        family: &'static str,
92        kind: ErrorKind,
93        #[source]
94        source: BoxError,
95    },
96    #[error("backend session entry failed: {source}")]
97    SessionEntry {
98        #[source]
99        source: crate::SessionEntryError,
100    },
101    #[error("missing runtime value for slot {slot}")]
102    MissingValue { slot: usize },
103    #[error("internal tensor error: {0}")]
104    Internal(String),
105}
106
107/// Owns the original tensor when a consuming representation reinterpretation
108/// cannot publish its checked descriptor.
109///
110/// Reinterpretation never falls back to an allocation or a copy.  Call
111/// [`Self::into_owner`] to recover the unchanged input and [`Self::error`] to
112/// inspect the typed failure.
113///
114/// # Examples
115///
116/// ```
117/// use tenferro_tensor::TypedTensor;
118///
119/// let tensor = TypedTensor::<f32>::from_vec_col_major(vec![1], vec![1.0])?;
120/// let Err(failure) = tensor.into_complex() else { return Ok(()); };
121/// assert!(!failure.error().to_string().is_empty());
122/// # Ok::<(), tenferro_tensor::Error>(())
123/// ```
124#[derive(Debug)]
125pub struct ReinterpretError<T> {
126    owner: Box<T>,
127    error: Error,
128}
129
130impl<T> ReinterpretError<T> {
131    pub(crate) fn new(owner: T, error: Error) -> Self {
132        Self {
133            owner: Box::new(owner),
134            error,
135        }
136    }
137
138    /// Recover the unchanged original owner.
139    ///
140    /// # Examples
141    ///
142    /// ```
143    /// use tenferro_tensor::TypedTensor;
144    ///
145    /// let tensor = TypedTensor::<f32>::from_vec_col_major(vec![1], vec![1.0])?;
146    /// let Err(failure) = tensor.into_complex() else { return Ok(()); };
147    /// let _owner = failure.into_owner();
148    /// # Ok::<(), tenferro_tensor::Error>(())
149    /// ```
150    pub fn into_owner(self) -> T {
151        *self.owner
152    }
153
154    /// Borrow the typed failure without consuming the owner.
155    ///
156    /// # Examples
157    ///
158    /// ```
159    /// use tenferro_tensor::TypedTensor;
160    ///
161    /// let tensor = TypedTensor::<f32>::from_vec_col_major(vec![1], vec![1.0])?;
162    /// let Err(failure) = tensor.into_complex() else { return Ok(()); };
163    /// assert!(!failure.error().to_string().is_empty());
164    /// # Ok::<(), tenferro_tensor::Error>(())
165    /// ```
166    pub fn error(&self) -> &Error {
167        &self.error
168    }
169
170    /// Consume the failure and return the retained owner with its typed cause.
171    ///
172    /// # Examples
173    ///
174    /// ```
175    /// use tenferro_tensor::TypedTensor;
176    ///
177    /// let tensor = TypedTensor::<f32>::from_vec_col_major(vec![1], vec![1.0])?;
178    /// let Err(failure) = tensor.into_complex() else { return Ok(()); };
179    /// let (owner, error) = failure.into_parts();
180    /// assert!(!error.to_string().is_empty());
181    /// assert_eq!(owner.shape(), &[1]);
182    /// # Ok::<(), tenferro_tensor::Error>(())
183    /// ```
184    #[must_use]
185    pub fn into_parts(self) -> (T, Error) {
186        (*self.owner, self.error)
187    }
188}
189
190impl<T: std::fmt::Debug> std::fmt::Display for ReinterpretError<T> {
191    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
192        write!(formatter, "tensor reinterpretation failed: {}", self.error)
193    }
194}
195
196impl<T: std::fmt::Debug + 'static> std::error::Error for ReinterpretError<T> {}
197
198impl Error {
199    /// Construct an incompatible-shapes validation error.
200    ///
201    /// # Examples
202    ///
203    /// ```rust
204    /// use tenferro_tensor::Error;
205    ///
206    /// let error = Error::shape_mismatch("add", [2, 3], [2, 4]);
207    /// assert!(matches!(error, Error::Validation { .. }));
208    /// ```
209    pub fn shape_mismatch(
210        op: &'static str,
211        lhs: impl Into<Vec<usize>>,
212        rhs: impl Into<Vec<usize>>,
213    ) -> Self {
214        Self::validation(
215            op,
216            tenferro_tensor_core::ShapeMismatch::IncompatibleShapes {
217                lhs: tenferro_tensor_core::ShapeVec::from_vec(lhs.into()),
218                rhs: tenferro_tensor_core::ShapeVec::from_vec(rhs.into()),
219            }
220            .into(),
221        )
222    }
223
224    /// Construct a rank-mismatch validation error.
225    ///
226    /// # Examples
227    ///
228    /// ```rust
229    /// use tenferro_tensor::Error;
230    ///
231    /// let error = Error::rank_mismatch("transpose", 2, 3);
232    /// assert!(matches!(error, Error::Validation { .. }));
233    /// ```
234    pub fn rank_mismatch(op: &'static str, expected: usize, actual: usize) -> Self {
235        Self::validation(op, ValidationError::RankMismatch { expected, actual })
236    }
237
238    /// Construct an axis-out-of-bounds validation error.
239    ///
240    /// # Examples
241    ///
242    /// ```rust
243    /// use tenferro_tensor::Error;
244    ///
245    /// let error = Error::axis_out_of_bounds("sum", 2, 2);
246    /// assert!(matches!(error, Error::Validation { .. }));
247    /// ```
248    pub fn axis_out_of_bounds(op: &'static str, axis: usize, rank: usize) -> Self {
249        Self::validation(op, ValidationError::AxisOutOfBounds { axis, rank })
250    }
251
252    /// Construct a duplicate-axis validation error.
253    ///
254    /// # Examples
255    ///
256    /// ```rust
257    /// use tenferro_tensor::Error;
258    ///
259    /// let error = Error::duplicate_axis("transpose", 1, "permutation");
260    /// assert!(matches!(error, Error::Validation { .. }));
261    /// ```
262    pub fn duplicate_axis(op: &'static str, axis: usize, role: &'static str) -> Self {
263        Self::validation(op, ValidationError::DuplicateAxis { axis, role })
264    }
265
266    /// Construct a dtype-mismatch validation error.
267    ///
268    /// # Examples
269    ///
270    /// ```rust
271    /// use tenferro_tensor::{DType, Error};
272    ///
273    /// let error = Error::dtype_mismatch("add", DType::F32, DType::F64);
274    /// assert!(matches!(error, Error::Validation { .. }));
275    /// ```
276    pub fn dtype_mismatch(op: &'static str, expected: crate::DType, actual: crate::DType) -> Self {
277        Self::validation(op, ValidationError::DTypeMismatch { expected, actual })
278    }
279
280    /// Wrap shared tensor validation with the operation that requested it.
281    ///
282    /// # Examples
283    ///
284    /// ```rust
285    /// use tenferro_tensor::{Error, ValidationError};
286    ///
287    /// let error = Error::validation(
288    ///     "transpose",
289    ///     ValidationError::AxisOutOfBounds { axis: 2, rank: 2 },
290    /// );
291    /// assert!(matches!(error, Error::Validation { op: "transpose", .. }));
292    /// ```
293    pub fn validation(op: &'static str, source: ValidationError) -> Self {
294        Self::Validation { op, source }
295    }
296
297    /// Construct a structured invalid-argument validation error.
298    ///
299    /// # Examples
300    ///
301    /// ```rust
302    /// use tenferro_tensor::{Error, ErrorKind, ValidationKind};
303    ///
304    /// let error = Error::invalid_argument("slice", "step", "must be non-zero");
305    /// assert_eq!(error.kind(), ErrorKind::Validation(ValidationKind::InvalidArgument));
306    /// ```
307    pub fn invalid_argument(
308        op: &'static str,
309        argument: &'static str,
310        message: impl Into<String>,
311    ) -> Self {
312        Self::validation(
313            op,
314            ValidationError::InvalidArgument {
315                argument,
316                message: message.into(),
317            },
318        )
319    }
320
321    /// Construct an unsupported dtype conversion error.
322    ///
323    /// # Examples
324    ///
325    /// ```rust
326    /// let error = tenferro_tensor::Error::unsupported_dtype_conversion(
327    ///     "convert",
328    ///     tenferro_tensor::DType::F64,
329    ///     tenferro_tensor::DType::I32,
330    ///     "lossy conversion is disabled",
331    /// );
332    /// assert!(matches!(
333    ///     error,
334    ///     tenferro_tensor::Error::UnsupportedDTypeConversion { .. }
335    /// ));
336    /// ```
337    pub fn unsupported_dtype_conversion(
338        op: &'static str,
339        from: crate::DType,
340        to: crate::DType,
341        message: impl Into<String>,
342    ) -> Self {
343        Self::UnsupportedDTypeConversion {
344            op,
345            from,
346            to,
347            message: message.into(),
348        }
349    }
350
351    /// Construct an operation-level unsupported-dtype error.
352    ///
353    /// This is for an operation that cannot run for the supplied dtype. It is
354    /// deliberately distinct from [`Error::unsupported_dtype_conversion`],
355    /// which is reserved for an actual from-dtype to to-dtype conversion.
356    ///
357    /// # Examples
358    ///
359    /// ```rust
360    /// let error = tenferro_tensor::Error::unsupported_dtype(
361    ///     "exp",
362    ///     tenferro_tensor::DType::I64,
363    ///     "integer exponentials are not implemented",
364    /// );
365    /// assert!(matches!(
366    ///     error,
367    ///     tenferro_tensor::Error::UnsupportedDType {
368    ///         op: "exp",
369    ///         dtype: tenferro_tensor::DType::I64,
370    ///         ..
371    ///     }
372    /// ));
373    /// ```
374    pub fn unsupported_dtype(
375        op: &'static str,
376        dtype: crate::DType,
377        message: impl Into<String>,
378    ) -> Self {
379        Self::UnsupportedDType {
380            op,
381            dtype,
382            message: message.into(),
383        }
384    }
385
386    /// Construct a structured unsupported-operation error.
387    ///
388    /// Use this for an operation or execution surface that is not implemented
389    /// by the selected backend. Dtype conversion failures use
390    /// [`Error::unsupported_dtype_conversion`] instead, and operation-specific
391    /// typed reasons should use [`Error::extension`] with `ErrorKind::Unsupported`.
392    ///
393    /// # Examples
394    ///
395    /// ```rust
396    /// let error = tenferro_tensor::Error::unsupported(
397    ///     "full_piv_lu",
398    ///     "backend has no implementation",
399    /// );
400    /// assert!(matches!(
401    ///     error,
402    ///     tenferro_tensor::Error::Unsupported { op: "full_piv_lu", .. }
403    /// ));
404    /// ```
405    pub fn unsupported(op: &'static str, message: impl Into<String>) -> Self {
406        Self::Unsupported {
407            op,
408            message: message.into(),
409        }
410    }
411
412    /// Construct a text-only backend failure.
413    ///
414    /// Use [`Error::backend_source`] when a typed source is available.
415    ///
416    /// # Examples
417    ///
418    /// ```rust
419    /// let error = tenferro_tensor::Error::backend_failure(
420    ///     "matmul",
421    ///     "backend rejected launch",
422    /// );
423    /// assert!(matches!(
424    ///     error,
425    ///     tenferro_tensor::Error::BackendFailure { op: "matmul", .. }
426    /// ));
427    /// ```
428    pub fn backend_failure(op: &'static str, message: impl Into<String>) -> Self {
429        Self::BackendFailure {
430            op,
431            message: message.into(),
432        }
433    }
434
435    /// Construct a backend failure while preserving its typed source.
436    ///
437    /// # Examples
438    ///
439    /// ```rust
440    /// let error = tenferro_tensor::Error::backend_source(
441    ///     "load",
442    ///     std::io::Error::other("read failed"),
443    /// );
444    /// assert!(std::error::Error::source(&error).is_some());
445    /// ```
446    pub fn backend_source<E>(op: &'static str, source: E) -> Self
447    where
448        E: StdError + Send + Sync + 'static,
449    {
450        Self::BackendSource {
451            op,
452            source: Box::new(source),
453        }
454    }
455
456    /// Construct an I/O failure while preserving its typed source.
457    ///
458    /// I/O errors are intentionally separate from backend failures: callers
459    /// can classify them as [`ErrorKind::Io`] without parsing a message.
460    ///
461    /// # Examples
462    ///
463    /// ```rust
464    /// use tenferro_tensor::{Error, ErrorKind};
465    ///
466    /// let error = Error::io_source("load", std::io::Error::other("read failed"));
467    /// assert_eq!(error.kind(), ErrorKind::Io);
468    /// assert!(std::error::Error::source(&error).is_some());
469    /// ```
470    pub fn io_source<E>(op: &'static str, source: E) -> Self
471    where
472        E: StdError + Send + Sync + 'static,
473    {
474        Self::IoSource {
475            op,
476            source: Box::new(source),
477        }
478    }
479
480    /// Construct a runtime-state failure when no typed source exists.
481    ///
482    /// Use this for missing, uninitialized, or invalid execution state. It is
483    /// distinct from [`Error::backend_failure`], which is reserved for
484    /// vendor/backend status text.
485    ///
486    /// # Examples
487    ///
488    /// ```rust
489    /// use tenferro_tensor::{Error, ErrorKind};
490    ///
491    /// let error = Error::runtime_state("execute", "backend session is not initialized");
492    /// assert_eq!(error.kind(), ErrorKind::RuntimeState);
493    /// ```
494    pub fn runtime_state(op: &'static str, message: impl Into<String>) -> Self {
495        Self::RuntimeState {
496            op,
497            message: message.into(),
498        }
499    }
500
501    /// Construct a runtime-state failure while preserving a typed source.
502    ///
503    /// # Examples
504    ///
505    /// ```rust
506    /// use tenferro_tensor::{Error, ErrorKind};
507    ///
508    /// let error = Error::runtime_state_source(
509    ///     "execute",
510    ///     std::io::Error::other("executor lock poisoned"),
511    /// );
512    /// assert_eq!(error.kind(), ErrorKind::RuntimeState);
513    /// assert!(std::error::Error::source(&error).is_some());
514    /// ```
515    pub fn runtime_state_source<E>(op: &'static str, source: E) -> Self
516    where
517        E: StdError + Send + Sync + 'static,
518    {
519        Self::RuntimeStateSource {
520            op,
521            source: Box::new(source),
522        }
523    }
524
525    /// Construct an extension failure while preserving its typed source and
526    /// coarse classification.
527    ///
528    /// # Examples
529    ///
530    /// ```rust
531    /// use std::error::Error as _;
532    /// use tenferro_tensor::{Error, ErrorKind};
533    ///
534    /// let error = Error::extension(
535    ///     "einsum",
536    ///     "einsum",
537    ///     ErrorKind::Internal,
538    ///     std::io::Error::other("planner failed"),
539    /// );
540    /// assert!(error.source().is_some());
541    /// ```
542    pub fn extension<E>(op: &'static str, family: &'static str, kind: ErrorKind, source: E) -> Self
543    where
544        E: StdError + Send + Sync + 'static,
545    {
546        Self::Extension {
547            op,
548            family,
549            kind,
550            source: Box::new(source),
551        }
552    }
553
554    /// Preserve a typed guarded-host-access failure.
555    ///
556    /// # Examples
557    ///
558    /// ```rust
559    /// use tenferro_tensor::{Error, HostAccessError};
560    ///
561    /// let error = Error::host_access(
562    ///     "map",
563    ///     HostAccessError::Unsupported { backend: "opaque" },
564    /// );
565    /// assert!(matches!(error, Error::HostAccess { .. }));
566    /// ```
567    pub fn host_access(op: &'static str, source: crate::HostAccessError) -> Self {
568        Self::HostAccess { op, source }
569    }
570
571    /// Return the stable coarse classification for this tensor failure.
572    ///
573    /// # Examples
574    ///
575    /// ```rust
576    /// use tenferro_tensor::{Error, ErrorKind, ValidationError, ValidationKind};
577    /// use tenferro_tensor::core::DType;
578    ///
579    /// let error = Error::validation(
580    ///     "add",
581    ///     ValidationError::DTypeMismatch {
582    ///         expected: DType::F32,
583    ///         actual: DType::F64,
584    ///     },
585    /// );
586    /// assert_eq!(error.kind(), ErrorKind::Validation(ValidationKind::DTypeMismatch));
587    /// ```
588    pub fn kind(&self) -> ErrorKind {
589        match self {
590            Self::Validation { source, .. } => ErrorKind::Validation(source.kind()),
591            Self::UnsupportedDTypeConversion { .. }
592            | Self::UnsupportedDType { .. }
593            | Self::Unsupported { .. } => ErrorKind::Unsupported,
594            Self::BackendFailure { .. } | Self::BackendSource { .. } => ErrorKind::BackendFailure,
595            Self::IoSource { .. } => ErrorKind::Io,
596            Self::RuntimeState { .. }
597            | Self::RuntimeStateSource { .. }
598            | Self::HostAccess { .. }
599            | Self::SessionEntry { .. } => ErrorKind::RuntimeState,
600            Self::Extension { kind, .. } => *kind,
601            Self::MissingValue { .. } => ErrorKind::RuntimeState,
602            Self::Internal(_) => ErrorKind::Internal,
603        }
604    }
605}
606
607impl From<crate::SessionEntryError> for Error {
608    fn from(source: crate::SessionEntryError) -> Self {
609        Self::SessionEntry { source }
610    }
611}
612
613/// Result type alias for runtime tensor operations.
614pub type Result<T> = std::result::Result<T, Error>;