use std::path::PathBuf;
use objc2_foundation::NSError;
use crate::DataType;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[error("{domain} (code {code}): {message}")]
pub struct NsErrorInfo {
domain: String,
code: isize,
message: String,
}
impl NsErrorInfo {
pub(crate) fn from_ns_error(error: &NSError) -> Self {
let (domain, code, message) = (
error.domain().to_string(),
error.code(),
error.localizedDescription().to_string(),
);
Self {
domain,
code,
message,
}
}
#[inline(always)]
pub fn domain(&self) -> &str {
&self.domain
}
#[inline(always)]
pub const fn code(&self) -> isize {
self.code
}
#[inline(always)]
pub fn message(&self) -> &str {
&self.message
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum LoadError {
#[error("model not found at `{}`", .0.display())]
NotFound(PathBuf),
#[error("core ml failed to load model: {0}")]
Native(NsErrorInfo),
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum CompileError {
#[error("model source not found at `{}`", .0.display())]
NotFound(PathBuf),
#[error("core ml failed to compile model: {0}")]
Native(NsErrorInfo),
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum PredictionError {
#[error("prediction output is missing feature `{0}`")]
MissingOutput(String),
#[error("prediction output `{0}` is not a multi-array")]
NotMultiArray(String),
#[error("core ml prediction failed: {0}")]
Native(NsErrorInfo),
#[error("stateful prediction is unavailable on this OS (requires macOS 15)")]
StateUnsupported,
#[error("failed to de-alias a prediction output: {0}")]
AliasCopyFailed(TensorError),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DataTypeMismatch {
expected: DataType,
actual: DataType,
}
impl DataTypeMismatch {
#[inline(always)]
pub const fn new(expected: DataType, actual: DataType) -> Self {
Self { expected, actual }
}
#[inline(always)]
pub const fn expected(&self) -> DataType {
self.expected
}
#[inline(always)]
pub const fn actual(&self) -> DataType {
self.actual
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ShapeMismatch {
expected: usize,
actual: usize,
}
impl ShapeMismatch {
#[inline(always)]
pub const fn new(expected: usize, actual: usize) -> Self {
Self { expected, actual }
}
#[inline(always)]
pub const fn expected(&self) -> usize {
self.expected
}
#[inline(always)]
pub const fn actual(&self) -> usize {
self.actual
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RankMismatch {
expected: usize,
actual: usize,
}
impl RankMismatch {
#[inline(always)]
pub const fn new(expected: usize, actual: usize) -> Self {
Self { expected, actual }
}
#[inline(always)]
pub const fn expected(&self) -> usize {
self.expected
}
#[inline(always)]
pub const fn actual(&self) -> usize {
self.actual
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct IndexOutOfBounds {
index: usize,
len: usize,
}
#[allow(clippy::len_without_is_empty)]
impl IndexOutOfBounds {
#[inline(always)]
pub const fn new(index: usize, len: usize) -> Self {
Self { index, len }
}
#[inline(always)]
pub const fn index(&self) -> usize {
self.index
}
#[inline(always)]
pub const fn len(&self) -> usize {
self.len
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NonContiguous {
shape: Vec<usize>,
strides: Vec<usize>,
}
impl NonContiguous {
#[inline(always)]
pub const fn new(shape: Vec<usize>, strides: Vec<usize>) -> Self {
Self { shape, strides }
}
#[inline(always)]
pub fn shape(&self) -> &[usize] {
&self.shape
}
#[inline(always)]
pub fn strides(&self) -> &[usize] {
&self.strides
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct UnsupportedShape {
shape: Vec<usize>,
reason: ShapeRequirement,
}
impl UnsupportedShape {
#[inline(always)]
pub const fn new(shape: Vec<usize>, reason: ShapeRequirement) -> Self {
Self { shape, reason }
}
#[inline(always)]
pub fn shape(&self) -> &[usize] {
&self.shape
}
#[inline(always)]
pub const fn reason(&self) -> ShapeRequirement {
self.reason
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct AllocationFailed {
elements: usize,
data_type: DataType,
}
impl AllocationFailed {
#[inline(always)]
pub const fn new(elements: usize, data_type: DataType) -> Self {
Self {
elements,
data_type,
}
}
#[inline(always)]
pub const fn elements(&self) -> usize {
self.elements
}
#[inline(always)]
pub const fn data_type(&self) -> DataType {
self.data_type
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum TensorError {
#[error("data type mismatch: expected `{}`, got `{}`", .0.expected(), .0.actual())]
DataTypeMismatch(DataTypeMismatch),
#[error("shape mismatch: expected {} elements, got {}", .0.expected(), .0.actual())]
ShapeMismatch(ShapeMismatch),
#[error("rank mismatch: expected {} indices, got {}", .0.expected(), .0.actual())]
RankMismatch(RankMismatch),
#[error("index {} out of bounds for length {}", .0.index(), .0.len())]
IndexOutOfBounds(IndexOutOfBounds),
#[error("array layout is not contiguous (strides {:?} for shape {:?})", .0.strides(), .0.shape())]
NonContiguous(NonContiguous),
#[error("unsupported data type `{0}` for array construction")]
UnsupportedDataType(DataType),
#[error("core ml multi-array failure: {0}")]
Native(NsErrorInfo),
#[error("pixel buffer creation failed with CVReturn {0}")]
PixelBuffer(i32),
#[error("shape {:?} is unsupported: {}", .0.shape(), .0.reason())]
UnsupportedShape(UnsupportedShape),
#[error("shape {0:?} element count overflows usize")]
ShapeOverflow(Vec<usize>),
#[error("pixel-buffer-backed arrays require macOS 12 or newer")]
SurfaceUnsupported,
#[error(
"could not allocate a {} element `{}` buffer for this array",
.0.elements(), .0.data_type()
)]
AllocationFailed(AllocationFailed),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, derive_more::Display, derive_more::IsVariant)]
#[display("{}", self.as_str())]
#[non_exhaustive]
pub enum ShapeRequirement {
LeadingDimsUnit,
NonEmpty,
NonZeroDims,
}
impl ShapeRequirement {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
Self::LeadingDimsUnit => "all dimensions before the last must be 1",
Self::NonEmpty => "the shape must have at least one dimension",
Self::NonZeroDims => "every dimension must be nonzero",
}
}
}
#[cfg(test)]
mod tests;