use indicatif::{ProgressBar, ProgressStyle};
use reqwest::blocking::Client;
use std::fs;
use std::io;
use std::path::{Path, PathBuf};
pub const MODEL_REPO: &str = "janakhpon/monocr";
pub const MODEL_REVISION: &str = "d3d9d5e";
pub const MODEL_FILENAME: &str = "monocr.onnx";
pub const CHARSET_FILENAME: &str = "charset.txt";
pub struct ModelManager {
cache_dir: PathBuf,
base_url: String,
model_filename: String,
}
impl Default for ModelManager {
fn default() -> Self {
Self::new()
}
}
impl ModelManager {
pub fn new() -> Self {
let home = dirs::home_dir().expect("Failed to get home directory");
let cache_dir = home.join(".monocr").join("models").join(MODEL_REVISION);
Self {
cache_dir,
base_url: format!("https://huggingface.co/{MODEL_REPO}/resolve/{MODEL_REVISION}"),
model_filename: MODEL_FILENAME.to_string(),
}
}
pub fn cache_dir(&self) -> &Path {
&self.cache_dir
}
pub fn model_url(&self) -> String {
format!("{}/onnx/{}", self.base_url, self.model_filename)
}
pub fn charset_url(&self) -> String {
format!("{}/onnx/{}", self.base_url, CHARSET_FILENAME)
}
pub fn get_model_path(&self) -> io::Result<PathBuf> {
let model_path = self.cache_dir.join(&self.model_filename);
if !model_path.exists() {
println!(
"Model {MODEL_REVISION} not found at {:?}. Downloading...",
model_path
);
self.download(&self.model_url(), &model_path)?;
println!("Download complete");
}
Ok(model_path)
}
pub fn get_charset(&self) -> io::Result<String> {
let charset_path = self.cache_dir.join(CHARSET_FILENAME);
if !charset_path.exists() {
self.download(&self.charset_url(), &charset_path)?;
}
fs::read_to_string(&charset_path)
}
fn download(&self, url: &str, dest: &Path) -> io::Result<()> {
if let Some(parent) = dest.parent() {
fs::create_dir_all(parent)?;
}
let client = Client::new();
let mut response = client
.get(url)
.send()
.map_err(|e| io::Error::other(format!("Failed to download from {}: {}", url, e)))?;
if !response.status().is_success() {
return Err(io::Error::other(format!(
"Failed to download {}: {}",
url,
response.status()
)));
}
let total_size = response.content_length().unwrap_or(0);
let pb = ProgressBar::new(total_size);
pb.set_style(ProgressStyle::default_bar()
.template("{spinner:.green} [{elapsed_precise}] [{wide_bar:.cyan/blue}] {bytes}/{total_bytes} ({eta})")
.unwrap()
.progress_chars(">-"));
let tmp_path = dest.with_extension("part");
let mut file = fs::File::create(&tmp_path)?;
let copied = io::copy(&mut response, &mut file);
drop(file);
if let Err(e) = copied {
let _ = fs::remove_file(&tmp_path);
return Err(e);
}
fs::rename(&tmp_path, dest)?;
pb.finish_and_clear();
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn download_urls_are_pinned() {
let m = ModelManager::new();
for url in [m.model_url(), m.charset_url()] {
assert!(
!url.contains("/resolve/main/"),
"still tracking the moving ref `main`: {url}"
);
assert!(
url.contains(&format!("/resolve/{MODEL_REVISION}/")),
"not pinned to {MODEL_REVISION}: {url}"
);
}
}
#[test]
fn charset_is_fetched_from_the_model_revision() {
let m = ModelManager::new();
let model_dir = m.model_url().trim_end_matches(MODEL_FILENAME).to_string();
let charset_dir = m
.charset_url()
.trim_end_matches(CHARSET_FILENAME)
.to_string();
assert_eq!(model_dir, charset_dir);
}
#[test]
fn cache_dir_is_scoped_by_revision() {
let m = ModelManager::new();
assert_eq!(
m.cache_dir().file_name().and_then(|s| s.to_str()),
Some(MODEL_REVISION),
"cache directory {:?} is not revision-scoped",
m.cache_dir()
);
}
}