use std::path::PathBuf;
use serde::{Deserialize, Serialize};
use crate::error::{OcrError, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Language {
#[default]
English,
Latin,
ChineseSimplified,
Japanese,
Korean,
Cyrillic,
Telugu,
Kannada,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Backend {
#[default]
Ort,
Tract,
Candle,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct ModelConfig {
pub languages: Vec<Language>,
pub backend: Backend,
pub detector_path: Option<PathBuf>,
pub recognizer_path: Option<PathBuf>,
pub cache_dir: Option<PathBuf>,
pub registry_owner: Option<String>,
}
impl Default for ModelConfig {
fn default() -> Self {
Self {
languages: vec![Language::English],
backend: Backend::default(),
detector_path: None,
recognizer_path: None,
cache_dir: None,
registry_owner: None,
}
}
}
impl ModelConfig {
pub(crate) fn validate(&self) -> Result<()> {
match (&self.detector_path, &self.recognizer_path) {
(Some(_), None) => Err(OcrError::config(
"model.recognizer_path is required when model.detector_path is configured",
)),
(None, Some(_)) => Err(OcrError::config(
"model.detector_path is required when model.recognizer_path is configured",
)),
_ => Ok(()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn should_accept_model_paths_only_as_a_pair() {
let mut config = ModelConfig {
detector_path: Some("craft.onnx".into()),
..ModelConfig::default()
};
assert!(config.validate().is_err());
config.recognizer_path = Some("english.onnx".into());
config.validate().expect("paired paths are valid");
}
}