准备 - core:内置官方字符集(OLD/BETA)与 ModelMetadata::from_builtin_* 构造器 - ort:导出 Session、实现 Info、修正 cuda feature 接线、共享推理工具 - tract2:包名更名(原 ddddocr-tract 已被占用)并完成规范整改 - 集成测试按领域拆分(ocr / det / slide / common / api_surface) - 发布准备:Cargo.toml 元数据、workspace 版本 0.2.4、LICENSE/ NOTICE、README 模型下载说明
102 lines
3.6 KiB
Rust
102 lines
3.6 KiB
Rust
use ort::Error as OrtError;
|
||
/// 模型加载与解析的通用结果类型。
|
||
pub type Result<T> = std::result::Result<T, Error>;
|
||
#[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<dyn std::error::Error + Send + Sync>),
|
||
}
|
||
|
||
impl Error {
|
||
/// 方便将任何第三方 Error 包装为 Error::Other
|
||
pub fn new<E>(msg: impl Into<String>, err: E) -> Self
|
||
where
|
||
E: Into<Box<dyn std::error::Error + Send + Sync>>,
|
||
{
|
||
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::<serde_json::Value>("{").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(_)));
|
||
}
|
||
}
|