use clap::Args;
use sceptre::OcrConfig;
const MIN_PROBABILITY: f32 = 0.0;
const MAX_PROBABILITY: f32 = 1.0;
fn parse_probability(raw: &str) -> core::result::Result<f32, String> {
let value: f32 = raw.parse().map_err(|_| format!("`{raw}` is not a number"))?;
if value.is_nan() || !(MIN_PROBABILITY..=MAX_PROBABILITY).contains(&value) {
return Err(format!("must be between {MIN_PROBABILITY} and {MAX_PROBABILITY}"));
}
Ok(value)
}
#[derive(Debug, Default, Args)]
pub struct OcrOverrides {
#[arg(long = "lang", value_enum)]
languages: Vec<LanguageArg>,
#[arg(long)]
threads: Option<usize>,
#[arg(long, value_enum)]
backend: Option<BackendArg>,
#[arg(long, value_parser = parse_probability)]
text_threshold: Option<f32>,
#[arg(long, value_parser = parse_probability)]
link_threshold: Option<f32>,
#[arg(long)]
canvas_size: Option<u32>,
}
#[derive(Debug, Clone, Copy, clap::ValueEnum)]
pub enum LanguageArg {
English,
Latin,
ChineseSimplified,
Japanese,
Korean,
Cyrillic,
Telugu,
Kannada,
}
#[derive(Debug, Clone, Copy, clap::ValueEnum)]
pub enum BackendArg {
Ort,
Tract,
Candle,
}
impl From<LanguageArg> for sceptre::Language {
fn from(value: LanguageArg) -> Self {
use sceptre::Language;
match value {
LanguageArg::English => Language::English,
LanguageArg::Latin => Language::Latin,
LanguageArg::ChineseSimplified => Language::ChineseSimplified,
LanguageArg::Japanese => Language::Japanese,
LanguageArg::Korean => Language::Korean,
LanguageArg::Cyrillic => Language::Cyrillic,
LanguageArg::Telugu => Language::Telugu,
LanguageArg::Kannada => Language::Kannada,
}
}
}
impl From<BackendArg> for sceptre::Backend {
fn from(value: BackendArg) -> Self {
use sceptre::Backend;
match value {
BackendArg::Ort => Backend::Ort,
BackendArg::Tract => Backend::Tract,
BackendArg::Candle => Backend::Candle,
}
}
}
impl OcrOverrides {
pub fn apply(&self, config: &mut OcrConfig) {
if !self.languages.is_empty() {
config.model.languages = self.languages.iter().copied().map(Into::into).collect();
}
if let Some(threads) = self.threads {
config.concurrency.max_threads = Some(threads);
}
if let Some(backend) = self.backend {
config.model.backend = backend.into();
}
if let Some(text_threshold) = self.text_threshold {
config.detection.text_threshold = text_threshold;
}
if let Some(link_threshold) = self.link_threshold {
config.detection.link_threshold = link_threshold;
}
if let Some(canvas_size) = self.canvas_size {
config.detection.canvas_size = canvas_size;
}
}
}
#[cfg(test)]
mod tests {
use super::parse_probability;
#[test]
fn should_accept_probabilities_within_the_unit_interval() {
assert_eq!(parse_probability("0.0").expect("zero is valid"), 0.0);
assert_eq!(parse_probability("1.0").expect("one is valid"), 1.0);
assert_eq!(parse_probability("0.5").expect("a mid value is valid"), 0.5);
}
#[test]
fn should_reject_out_of_range_nan_and_unparseable_probabilities() {
assert!(parse_probability("-0.1").is_err(), "a negative value is rejected");
assert!(parse_probability("1.1").is_err(), "a value above one is rejected");
assert!(parse_probability("nan").is_err(), "NaN is rejected");
assert!(parse_probability("abc").is_err(), "a non-numeric value is rejected");
}
}