1use tenferro_tensor::{DType, ErrorKind};
22
23#[cfg(any(feature = "cpu-blas", feature = "cuda", feature = "cpu-faer"))]
25#[derive(Debug, thiserror::Error)]
26pub(crate) enum BackendError {
27 #[cfg(feature = "cuda")]
28 #[error("{library} call {call} returned status {status}")]
29 ProviderStatus {
30 library: &'static str,
31 call: &'static str,
32 status: i32,
33 },
34 #[cfg(any(feature = "cpu-blas", feature = "cpu-faer"))]
35 #[error("{library} routine {routine} returned an invalid workspace: {detail}")]
36 InvalidWorkspace {
37 library: &'static str,
38 routine: &'static str,
39 detail: String,
40 },
41}
42
43#[derive(Debug, thiserror::Error)]
55#[non_exhaustive]
56pub enum Error {
57 #[error("{op} did not converge")]
59 NonConvergence {
60 op: &'static str,
62 },
63 #[error("{op} encountered non-finite {role}")]
65 NonFinite {
66 op: &'static str,
68 role: &'static str,
70 },
71 #[error("{op} is singular")]
73 Singular {
74 op: &'static str,
76 },
77 #[error("{op} does not support dtype {dtype:?}")]
79 UnsupportedDType {
80 op: &'static str,
82 dtype: DType,
84 },
85}
86
87impl Error {
88 #[must_use]
102 pub fn kind(&self) -> ErrorKind {
103 match self {
104 Self::NonConvergence { .. } | Self::NonFinite { .. } | Self::Singular { .. } => {
105 ErrorKind::NumericalFailure
106 }
107 Self::UnsupportedDType { .. } => ErrorKind::Unsupported,
108 }
109 }
110}
111
112pub type Result<T> = std::result::Result<T, Error>;
123
124pub(crate) fn into_tensor_error(op: &'static str, source: Error) -> tenferro_tensor::Error {
126 tenferro_tensor::Error::extension(
127 op,
128 crate::extension::LINALG_EXTENSION_FAMILY_ID,
129 source.kind(),
130 source,
131 )
132}
133
134pub(crate) fn unsupported_dtype(op: &'static str, dtype: DType) -> tenferro_tensor::Error {
136 into_tensor_error(op, Error::UnsupportedDType { op, dtype })
137}
138
139#[cfg(feature = "cuda")]
141pub(crate) fn backend_status(
142 op: &'static str,
143 library: &'static str,
144 call: &'static str,
145 status: i32,
146) -> tenferro_tensor::Error {
147 tenferro_tensor::Error::backend_source(
148 op,
149 BackendError::ProviderStatus {
150 library,
151 call,
152 status,
153 },
154 )
155}
156
157#[cfg(any(feature = "cpu-blas", feature = "cpu-faer"))]
162pub(crate) fn invalid_workspace(
163 op: &'static str,
164 library: &'static str,
165 routine: &'static str,
166 detail: impl Into<String>,
167) -> tenferro_tensor::Error {
168 tenferro_tensor::Error::backend_source(
169 op,
170 BackendError::InvalidWorkspace {
171 library,
172 routine,
173 detail: detail.into(),
174 },
175 )
176}
177
178#[cfg(all(
179 test,
180 any(feature = "cpu-blas", feature = "cuda", feature = "cpu-faer")
181))]
182mod tests {
183 use std::error::Error as _;
184
185 use super::*;
186
187 #[cfg(feature = "cuda")]
188 #[test]
189 fn provider_status_keeps_typed_backend_source() {
190 let error = backend_status("svd", "cuSOLVER", "cusolverDnSgesvd", 7);
191
192 assert_eq!(error.kind(), ErrorKind::BackendFailure);
193 assert!(matches!(
194 error.source().and_then(|source| source.downcast_ref()),
195 Some(BackendError::ProviderStatus {
196 library: "cuSOLVER",
197 call: "cusolverDnSgesvd",
198 status: 7,
199 })
200 ));
201 }
202
203 #[cfg(any(feature = "cpu-blas", feature = "cpu-faer"))]
204 #[test]
205 fn invalid_workspace_keeps_typed_backend_source() {
206 let error = invalid_workspace("eigh", "LAPACK", "dsyevd", "query was zero");
207
208 assert_eq!(error.kind(), ErrorKind::BackendFailure);
209 assert!(matches!(
210 error.source().and_then(|source| source.downcast_ref()),
211 Some(BackendError::InvalidWorkspace {
212 library: "LAPACK",
213 routine: "dsyevd",
214 detail,
215 }) if detail == "query was zero"
216 ));
217 }
218}