use ort::Error as OrtError; /// 模型加载与解析的通用结果类型。 pub type Result = std::result::Result; #[derive(thiserror::Error, Debug)] /// 模型加载、解析与 Session 构建阶段的错误。 pub enum Error { /// 底层构建器配置失败。 #[error("builder构建失败")] Build(#[from] BuildError), /// 解析 ONNX 模型/路径失败(如文件损坏、算子不支持、路径非法) #[error("解析 ONNX 模型结构失败: {0}")] ModelParse(#[from] ParseError), /// 模型计算图优化失败(如常量折叠、形状推导失败) #[error("优化 ORT 模型图失败: {0}")] OptimizationFailed(#[source] OrtError), /// 构建可执行 Session 失败(如输入输出 Tensor 类型/形状未确定) #[error("构建可运行的 ORT Session 失败: {0}")] RunnableBuildFailed(#[source] OrtError), /// JSON 反序列化失败(自动透传 serde_json 报错) #[error("模型 Metadata JSON 解析失败: {0}")] JsonParse(#[from] serde_json::Error), /// 字节流非合法 UTF-8 编码(自动透传 Utf8Error) #[error("Metadata 字节流不是合法的 UTF-8 编码: {0}")] InvalidUtf8(#[from] std::str::Utf8Error), /// 模型元数据内容解析失败。 #[error("模型元数据解析失败: {0}")] MetadataParse(String), /// 承载任何第三方扩展、解密、特定预处理插件在执行时产生的自定义错误 #[error("{0}: {1}")] Other(String, #[source] Box), } impl Error { /// 方便将任何第三方 Error 包装为 Error::Other pub fn new(msg: impl Into, err: E) -> Self where E: Into>, { Self::Other(msg.into(), err.into()) } } #[derive(thiserror::Error, Debug)] /// 从路径或字节流解析 ONNX 模型失败的错误。 pub enum ParseError { /// 策略 A:从文件路径加载失败(附带路径上下文信息,方便排查是找不到文件还是格式不对) #[error("从路径 '{0}' 加载 ONNX 模型失败: {1}")] Path(String, #[source] OrtError), /// 策略 B:从内存字节流加载失败(如 include_bytes! 传入的字节流损坏) #[error("从内存字节流解析 ONNX 模型失败: {0}")] Bytes(#[source] OrtError), } #[derive(thiserror::Error, Debug)] pub enum BuildError { /// 底层 SessionBuilder 构建失败。 #[error("构建 ORT SessionBuilder 失败: {0}")] BuildFailed(#[from] OrtError), /// 线程数配置失败。 #[error("{0}")] Threads(String), /// 启用 CUDA 执行提供者失败。 #[error("{0}")] EnabledCudaFailed(String), /// 未编译 CUDA 支持时尝试启用 GPU。 #[error("{0}")] NotEnabledCuda(String), } #[cfg(test)] mod tests { use super::*; #[test] fn wraps_external_error_via_new() { let io_err = std::io::Error::other("boom"); let err = Error::new("自定义错误", io_err); assert!(matches!(err, Error::Other(_, _))); assert!(err.to_string().contains("自定义错误")); } #[test] fn serde_json_error_converts_to_error() { let json_err = serde_json::from_str::("{").unwrap_err(); let err: Error = json_err.into(); assert!(matches!(err, Error::JsonParse(_))); } #[test] fn ut8_error_converts_to_error() { let bytes = vec![0xffu8]; let utf8_err = std::str::from_utf8(&bytes).unwrap_err(); let err: Error = utf8_err.into(); assert!(matches!(err, Error::InvalidUtf8(_))); } }