use std::path::PathBuf;
use std::sync::Arc;
use crate::config::Language;
use crate::error::{OcrError, Result};
use crate::models::integrity::verify_sha256_bytes;
use crate::models::provision::{ModelDescriptor, ModelRole};
#[derive(Debug, Clone)]
pub enum ModelArtifact {
Path(PathBuf),
Bytes(Arc<[u8]>),
}
pub trait ModelProvider: Send + Sync {
fn detector(&self) -> Result<ModelArtifact>;
fn recognizer(&self, language: Language) -> Result<ModelArtifact>;
}
#[derive(Debug)]
pub struct VerifiedModelProvider {
detector: Arc<[u8]>,
recognizers: Vec<(Language, Arc<[u8]>)>,
}
impl VerifiedModelProvider {
pub fn new(artifacts: impl IntoIterator<Item = (ModelDescriptor, Arc<[u8]>)>) -> Result<Self> {
let mut detector = None;
let mut recognizers = Vec::new();
for (descriptor, bytes) in artifacts {
verify_sha256_bytes(&bytes, &descriptor.sha256, &descriptor.name)?;
match descriptor.role {
ModelRole::Detector => {
if detector.replace(bytes).is_some() {
return Err(OcrError::model(
"verified model provider received more than one detector",
));
}
}
ModelRole::Recognizer(language) => {
if recognizers.iter().any(|(existing, _)| *existing == language) {
return Err(OcrError::model(format!(
"verified model provider received more than one {language:?} recognizer"
)));
}
recognizers.push((language, bytes));
}
}
}
let detector = detector.ok_or_else(|| OcrError::model("verified model provider requires one detector"))?;
Ok(Self { detector, recognizers })
}
}
impl ModelProvider for VerifiedModelProvider {
fn detector(&self) -> Result<ModelArtifact> {
Ok(ModelArtifact::Bytes(self.detector.clone()))
}
fn recognizer(&self, language: Language) -> Result<ModelArtifact> {
self.recognizers
.iter()
.find(|(candidate, _)| *candidate == language)
.map(|(_, bytes)| ModelArtifact::Bytes(bytes.clone()))
.ok_or_else(|| OcrError::model(format!("no verified recognizer bytes were supplied for {language:?}")))
}
}