use std::path::{Path, PathBuf};
use crate::config::{Language, OcrConfig};
use crate::error::Result;
use super::download::{self, hf_cache_root, resolve_cached};
use super::registry::{ModelEntry, craft_entry, effective_repo, recognizer_entry};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ModelRole {
Detector,
Recognizer(Language),
}
#[derive(Debug, Clone)]
pub struct ModelInfo {
pub name: String,
pub repo: String,
pub role: ModelRole,
pub cached: bool,
pub path: Option<PathBuf>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ModelDescriptor {
pub name: String,
pub repo: String,
pub revision: String,
pub file: String,
pub sha256: String,
pub role: ModelRole,
}
pub fn model_descriptors(config: &OcrConfig) -> Result<Vec<ModelDescriptor>> {
let owner = config.model.registry_owner.as_deref();
let mut descriptors = Vec::with_capacity(config.model.languages.len() + 1);
descriptors.push(descriptor_for(&craft_entry(), ModelRole::Detector, owner)?);
let mut described_languages = Vec::new();
for &language in &config.model.languages {
if described_languages.contains(&language) {
continue;
}
described_languages.push(language);
descriptors.push(descriptor_for(
&recognizer_entry(language),
ModelRole::Recognizer(language),
owner,
)?);
}
Ok(descriptors)
}
pub fn model_manifest(config: &OcrConfig) -> Result<Vec<ModelInfo>> {
let cache_dir = resolve_cache_dir(config)?;
let owner = config.model.registry_owner.as_deref();
let mut manifest = Vec::with_capacity(config.model.languages.len() + 1);
manifest.push(info_for(&craft_entry(), ModelRole::Detector, &cache_dir, owner)?);
for &language in &config.model.languages {
let entry = recognizer_entry(language);
manifest.push(info_for(&entry, ModelRole::Recognizer(language), &cache_dir, owner)?);
}
Ok(manifest)
}
pub fn download_models(config: &OcrConfig) -> Result<Vec<ModelInfo>> {
let cache_override = config.model.cache_dir.as_deref();
let owner = config.model.registry_owner.as_deref();
let mut manifest = Vec::with_capacity(config.model.languages.len() + 1);
manifest.push(fetch(&craft_entry(), ModelRole::Detector, cache_override, owner)?);
for &language in &config.model.languages {
let entry = recognizer_entry(language);
manifest.push(fetch(&entry, ModelRole::Recognizer(language), cache_override, owner)?);
}
Ok(manifest)
}
fn resolve_cache_dir(config: &OcrConfig) -> Result<PathBuf> {
hf_cache_root(config.model.cache_dir.as_deref())
}
fn descriptor_for(entry: &ModelEntry, role: ModelRole, owner: Option<&str>) -> Result<ModelDescriptor> {
Ok(ModelDescriptor {
name: entry.name.to_string(),
repo: effective_repo(entry, owner)?,
revision: entry.revision.to_string(),
file: entry.file.to_string(),
sha256: entry.sha256.to_string(),
role,
})
}
fn info_for(entry: &ModelEntry, role: ModelRole, cache_dir: &Path, owner: Option<&str>) -> Result<ModelInfo> {
let repo = effective_repo(entry, owner)?;
let path = resolve_cached(cache_dir, &repo, entry.file);
let cached = path.is_some();
Ok(ModelInfo {
name: entry.name.to_string(),
repo,
role,
cached,
path,
})
}
fn fetch(entry: &ModelEntry, role: ModelRole, cache_override: Option<&Path>, owner: Option<&str>) -> Result<ModelInfo> {
let repo = effective_repo(entry, owner)?;
let path = download::ensure(entry, cache_override, owner)?;
Ok(ModelInfo {
name: entry.name.to_string(),
repo,
role,
cached: true,
path: Some(path),
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::OcrConfig;
const CRAFT_REPO: &str = "sceptre-ocr/craft_mlt_25k";
const ENGLISH_REPO: &str = "sceptre-ocr/english_g2";
fn config_with_empty_cache(languages: Vec<Language>) -> OcrConfig {
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let unique = COUNTER.fetch_add(1, Ordering::Relaxed);
let cache_dir = std::env::temp_dir().join(format!("sceptre-manifest-{}-{unique}", std::process::id()));
let mut config = OcrConfig::default();
config.model.languages = languages;
config.model.cache_dir = Some(cache_dir);
config
}
#[test]
fn model_manifest_lists_detector_then_english_recognizer_for_default_config() {
let config = config_with_empty_cache(vec![Language::English]);
let manifest = model_manifest(&config).unwrap();
assert_eq!(manifest.len(), 2);
assert_eq!(manifest[0].name, "craft_mlt_25k");
assert_eq!(manifest[0].role, ModelRole::Detector);
assert_eq!(manifest[0].repo, CRAFT_REPO);
assert_eq!(manifest[1].name, "english_g2");
assert_eq!(manifest[1].role, ModelRole::Recognizer(Language::English));
assert_eq!(manifest[1].repo, ENGLISH_REPO);
}
#[test]
fn model_manifest_reports_not_cached_against_a_fresh_temp_cache_dir() {
let config = config_with_empty_cache(vec![Language::English]);
let manifest = model_manifest(&config).unwrap();
for info in &manifest {
assert!(!info.cached, "expected `{}` to be uncached", info.name);
assert_eq!(info.path, None);
}
}
#[test]
#[cfg(not(feature = "download"))]
fn download_models_errors_without_the_download_feature() {
let config = config_with_empty_cache(vec![Language::English]);
let error = download_models(&config).expect_err("download requires the `download` feature");
assert!(
matches!(error, crate::error::OcrError::Model { .. }),
"expected OcrError::Model, got {error:?}"
);
}
#[test]
fn model_manifest_yields_one_recognizer_per_language_in_order() {
let languages = vec![Language::Cyrillic, Language::Japanese, Language::Korean];
let config = config_with_empty_cache(languages.clone());
let manifest = model_manifest(&config).unwrap();
assert_eq!(manifest.len(), languages.len() + 1);
assert_eq!(manifest[0].role, ModelRole::Detector);
let recognizer_roles: Vec<ModelRole> = manifest[1..].iter().map(|info| info.role.clone()).collect();
let expected: Vec<ModelRole> = languages.into_iter().map(ModelRole::Recognizer).collect();
assert_eq!(recognizer_roles, expected);
}
#[test]
fn model_descriptors_deduplicate_repeated_languages() {
let mut config = OcrConfig::default();
config.model.languages = vec![Language::English, Language::English, Language::Telugu];
let descriptors = model_descriptors(&config).expect("registry metadata should resolve");
assert_eq!(descriptors.len(), 3);
assert_eq!(descriptors[1].role, ModelRole::Recognizer(Language::English));
assert_eq!(descriptors[2].role, ModelRole::Recognizer(Language::Telugu));
}
}