use anyhow::{Context, Result};
use sha2::{Digest, Sha256};
use std::fs;
use std::io::{Read, Write};
use std::path::{Path, PathBuf};
use std::sync::Arc;
const MODEL_ONNX_URL: &str =
"https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2/resolve/c9745ed1d9f207416be6d2e6f8de32d1f16199bf/onnx/model.onnx";
const MODEL_QUANTIZED_URL: &str = "https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2/resolve/c9745ed1d9f207416be6d2e6f8de32d1f16199bf/onnx/model_quint8_avx2.onnx";
const TOKENIZER_URL: &str =
"https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2/resolve/c9745ed1d9f207416be6d2e6f8de32d1f16199bf/tokenizer.json";
const NER_MODEL_URL: &str =
"https://huggingface.co/onnx-community/TinyBERT-finetuned-NER-ONNX/resolve/9b03777d9832105fbe419f258127fb2ec3eb09d7/onnx/model_quantized.onnx";
const NER_TOKENIZER_URL: &str =
"https://huggingface.co/onnx-community/TinyBERT-finetuned-NER-ONNX/resolve/9b03777d9832105fbe419f258127fb2ec3eb09d7/tokenizer.json";
struct ModelChecksums;
impl ModelChecksums {
const QUANTIZED_MODEL: Option<&'static str> =
Some("b941bf19f1f1283680f449fa6a7336bb5600bdcd5f84d10ddc5cd72218a0fd21");
const FULL_MODEL: Option<&'static str> =
Some("6fd5d72fe4589f189f8ebc006442dbb529bb7ce38f8082112682524616046452");
const TOKENIZER: Option<&'static str> =
Some("be50c3628f2bf5bb5e3a7f17b1f74611b2561a3a27eeab05e5aa30f411572037");
const NER_MODEL: Option<&'static str> =
Some("ba4a1a00cf1600cae8e7cf3fda4650c825811719065b51041256392edd3647b8");
const NER_TOKENIZER: Option<&'static str> =
Some("d241a60d5e8f04cc1b2b3e9ef7a4921b27bf526d9f6050ab90f9267a1f9e5c66");
}
#[cfg(target_os = "windows")]
const ONNX_RUNTIME_URL: &str = "https://github.com/microsoft/onnxruntime/releases/download/v1.23.2/onnxruntime-win-x64-1.23.2.zip";
#[cfg(target_os = "linux")]
const ONNX_RUNTIME_URL: &str = "https://github.com/microsoft/onnxruntime/releases/download/v1.23.2/onnxruntime-linux-x64-1.23.2.tgz";
#[cfg(target_os = "macos")]
const ONNX_RUNTIME_URL: &str = "https://github.com/microsoft/onnxruntime/releases/download/v1.23.2/onnxruntime-osx-arm64-1.23.2.tgz";
pub fn get_cache_dir() -> PathBuf {
if let Some(cache) = dirs::cache_dir() {
return cache.join("shodh-memory");
}
if let Some(home) = dirs::home_dir() {
return home.join(".cache").join("shodh-memory");
}
PathBuf::from(".shodh-cache")
}
pub fn get_models_dir() -> PathBuf {
get_cache_dir().join("models").join("minilm-l6")
}
pub fn get_ner_models_dir() -> PathBuf {
get_cache_dir().join("models").join("bert-tiny-ner")
}
pub fn get_onnx_runtime_dir() -> PathBuf {
get_cache_dir().join("onnxruntime")
}
pub fn are_models_downloaded() -> bool {
let models_dir = get_models_dir();
let model_path = models_dir.join("model_quantized.onnx");
let tokenizer_path = models_dir.join("tokenizer.json");
if !model_path.exists() || !tokenizer_path.exists() {
return false;
}
if let Some(expected) = ModelChecksums::QUANTIZED_MODEL {
if let Ok(valid) = verify_checksum(&model_path, expected) {
if !valid {
tracing::warn!("Model file checksum mismatch — will re-download");
let _ = fs::remove_file(&model_path);
return false;
}
}
}
if let Some(expected) = ModelChecksums::TOKENIZER {
if let Ok(valid) = verify_checksum(&tokenizer_path, expected) {
if !valid {
tracing::warn!("Tokenizer file checksum mismatch — will re-download");
let _ = fs::remove_file(&tokenizer_path);
return false;
}
}
}
true
}
pub fn are_ner_models_downloaded() -> bool {
let models_dir = get_ner_models_dir();
let model_path = models_dir.join("model.onnx");
let tokenizer_path = models_dir.join("tokenizer.json");
if !model_path.exists() || !tokenizer_path.exists() {
return false;
}
if let Some(expected) = ModelChecksums::NER_MODEL {
if let Ok(valid) = verify_checksum(&model_path, expected) {
if !valid {
tracing::warn!("NER model file checksum mismatch — will re-download");
let _ = fs::remove_file(&model_path);
return false;
}
}
}
if let Some(expected) = ModelChecksums::NER_TOKENIZER {
if let Ok(valid) = verify_checksum(&tokenizer_path, expected) {
if !valid {
tracing::warn!("NER tokenizer file checksum mismatch — will re-download");
let _ = fs::remove_file(&tokenizer_path);
return false;
}
}
}
true
}
pub fn is_onnx_runtime_downloaded() -> bool {
let onnx_dir = get_onnx_runtime_dir();
#[cfg(target_os = "windows")]
let lib_name = "onnxruntime.dll";
#[cfg(target_os = "linux")]
let lib_name = "libonnxruntime.so";
#[cfg(target_os = "macos")]
let lib_name = "libonnxruntime.dylib";
let path = onnx_dir.join(lib_name);
if !path.exists() {
if path.symlink_metadata().is_ok() {
tracing::warn!(
"Removing dangling symlink at {:?} (target does not exist)",
path
);
let _ = fs::remove_file(&path);
}
return false;
}
true
}
pub fn get_onnx_runtime_path() -> Option<PathBuf> {
let onnx_dir = get_onnx_runtime_dir();
#[cfg(target_os = "windows")]
let lib_name = "onnxruntime.dll";
#[cfg(target_os = "linux")]
let lib_name = "libonnxruntime.so";
#[cfg(target_os = "macos")]
let lib_name = "libonnxruntime.dylib";
let path = onnx_dir.join(lib_name);
if path.exists() {
Some(path)
} else {
None
}
}
pub type ProgressCallback = Arc<dyn Fn(u64, u64) + Send + Sync>;
fn verify_checksum(path: &Path, expected: &str) -> Result<bool> {
let mut file = fs::File::open(path).context("Failed to open file for checksum")?;
let mut hasher = Sha256::new();
let mut buffer = [0u8; 8192];
loop {
let bytes_read = file.read(&mut buffer)?;
if bytes_read == 0 {
break;
}
hasher.update(&buffer[..bytes_read]);
}
let result = hasher.finalize();
let actual = hex::encode(result);
if actual == expected.to_lowercase() {
Ok(true)
} else {
tracing::warn!(
"Checksum mismatch for {:?}: expected {}, got {}",
path,
expected,
actual
);
Ok(false)
}
}
#[allow(dead_code)]
fn compute_checksum(path: &Path) -> Result<String> {
let mut file = fs::File::open(path).context("Failed to open file for checksum")?;
let mut hasher = Sha256::new();
let mut buffer = [0u8; 8192];
loop {
let bytes_read = file.read(&mut buffer)?;
if bytes_read == 0 {
break;
}
hasher.update(&buffer[..bytes_read]);
}
Ok(hex::encode(hasher.finalize()))
}
fn download_file(
url: &str,
path: &PathBuf,
progress: Option<&(dyn Fn(u64, u64) + Send + Sync)>,
) -> Result<()> {
download_file_with_checksum(url, path, progress, None)
}
fn download_file_with_checksum(
url: &str,
path: &PathBuf,
progress: Option<&(dyn Fn(u64, u64) + Send + Sync)>,
expected_checksum: Option<&str>,
) -> Result<()> {
tracing::info!("Downloading {} to {:?}", url, path);
if let Some(parent) = path.parent() {
fs::create_dir_all(parent).context("Failed to create cache directory")?;
}
let response = ureq::get(url)
.call()
.context(format!("Failed to download from {url}"))?;
let total_size = response
.headers()
.get("content-length")
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u64>().ok())
.unwrap_or(0);
let mut reader = response.into_body().into_reader();
let mut file = fs::File::create(path).context("Failed to create output file")?;
let mut hasher = Sha256::new();
let mut downloaded: u64 = 0;
let mut buffer = [0u8; 8192];
loop {
let bytes_read = reader
.read(&mut buffer)
.context("Failed to read from download stream")?;
if bytes_read == 0 {
break;
}
file.write_all(&buffer[..bytes_read])
.context("Failed to write to file")?;
hasher.update(&buffer[..bytes_read]);
downloaded += bytes_read as u64;
if let Some(cb) = progress {
cb(downloaded, total_size);
}
}
let actual_checksum = hex::encode(hasher.finalize());
tracing::info!(
"Downloaded {} bytes to {:?} (SHA-256: {})",
downloaded,
path,
actual_checksum
);
if let Some(expected) = expected_checksum {
if actual_checksum != expected.to_lowercase() {
if let Err(e) = fs::remove_file(path) {
tracing::error!("Failed to delete corrupted file {path:?}: {e}");
}
anyhow::bail!(
"Checksum verification failed for {path:?}. Expected: {expected}, Got: {actual_checksum}. File deleted for security."
);
}
tracing::info!("Checksum verified for {path:?}");
} else {
tracing::warn!(
"No checksum provided for {:?}. For security, add this checksum: {}",
path,
actual_checksum
);
}
Ok(())
}
pub fn download_models(progress: Option<ProgressCallback>) -> Result<PathBuf> {
download_models_internal(progress, true)
}
pub fn download_models_internal(
progress: Option<ProgressCallback>,
use_quantized: bool,
) -> Result<PathBuf> {
let models_dir = get_models_dir();
if are_models_downloaded() {
tracing::info!("Models already downloaded at {:?}", models_dir);
return Ok(models_dir);
}
tracing::info!("Downloading MiniLM-L6-v2 model to {:?}", models_dir);
let (model_url, model_filename, model_checksum) = if use_quantized {
(
MODEL_QUANTIZED_URL,
"model_quantized.onnx",
ModelChecksums::QUANTIZED_MODEL,
)
} else {
(MODEL_ONNX_URL, "model.onnx", ModelChecksums::FULL_MODEL)
};
let model_path = models_dir.join(model_filename);
tracing::info!(
"Downloading model from {} (~{}MB)",
if use_quantized {
"HuggingFace (quantized)"
} else {
"HuggingFace (full)"
},
if use_quantized { 23 } else { 90 }
);
download_file_with_checksum(
model_url,
&model_path,
progress.as_ref().map(|p| p.as_ref()),
model_checksum,
)?;
let tokenizer_path = models_dir.join("tokenizer.json");
tracing::info!("Downloading tokenizer.json");
download_file_with_checksum(
TOKENIZER_URL,
&tokenizer_path,
progress.as_ref().map(|p| p.as_ref()),
ModelChecksums::TOKENIZER,
)?;
tracing::info!(
"MiniLM-L6-v2 model downloaded successfully to {:?}",
models_dir
);
Ok(models_dir)
}
pub fn download_ner_models(progress: Option<ProgressCallback>) -> Result<PathBuf> {
let models_dir = get_ner_models_dir();
if are_ner_models_downloaded() {
tracing::info!("NER models already downloaded at {:?}", models_dir);
return Ok(models_dir);
}
tracing::info!(
"Downloading TinyBERT-NER model to {:?} (~14.5MB)",
models_dir
);
let model_path = models_dir.join("model.onnx");
tracing::info!("Downloading NER model_quantized.onnx (~14.5MB)");
download_file_with_checksum(
NER_MODEL_URL,
&model_path,
progress.as_ref().map(|p| p.as_ref()),
ModelChecksums::NER_MODEL,
)?;
let tokenizer_path = models_dir.join("tokenizer.json");
tracing::info!("Downloading NER tokenizer.json");
download_file_with_checksum(
NER_TOKENIZER_URL,
&tokenizer_path,
progress.as_ref().map(|p| p.as_ref()),
ModelChecksums::NER_TOKENIZER,
)?;
tracing::info!(
"TinyBERT-NER model downloaded successfully to {:?}",
models_dir
);
Ok(models_dir)
}
pub fn download_onnx_runtime(progress: Option<ProgressCallback>) -> Result<PathBuf> {
let onnx_dir = get_onnx_runtime_dir();
if is_onnx_runtime_downloaded() {
tracing::info!("ONNX Runtime already downloaded at {:?}", onnx_dir);
return get_onnx_runtime_path().ok_or_else(|| anyhow::anyhow!("ONNX Runtime not found"));
}
tracing::info!("Downloading ONNX Runtime to {:?}", onnx_dir);
fs::create_dir_all(&onnx_dir)?;
let archive_name = if cfg!(target_os = "windows") {
"onnxruntime.zip"
} else {
"onnxruntime.tgz"
};
let archive_path = onnx_dir.join(archive_name);
download_file(
ONNX_RUNTIME_URL,
&archive_path,
progress.as_ref().map(|p| p.as_ref()),
)?;
extract_onnx_runtime(&archive_path, &onnx_dir)?;
if let Err(e) = fs::remove_file(&archive_path) {
tracing::warn!("Failed to clean up archive {:?}: {}", archive_path, e);
}
get_onnx_runtime_path().ok_or_else(|| anyhow::anyhow!("Failed to extract ONNX Runtime"))
}
fn extract_onnx_runtime(archive_path: &Path, dest_dir: &Path) -> Result<()> {
tracing::info!("Extracting ONNX Runtime from {:?}", archive_path);
#[cfg(target_os = "windows")]
{
let file = fs::File::open(archive_path)?;
let mut archive = zip::ZipArchive::new(file)?;
for i in 0..archive.len() {
let mut file = archive.by_index(i)?;
let name = file.name();
if name.ends_with("onnxruntime.dll") {
let dest_path = dest_dir.join("onnxruntime.dll");
let mut outfile = fs::File::create(&dest_path)?;
std::io::copy(&mut file, &mut outfile)?;
tracing::info!("Extracted onnxruntime.dll");
return Ok(());
}
}
anyhow::bail!("onnxruntime.dll not found in archive");
}
#[cfg(not(target_os = "windows"))]
{
let file = fs::File::open(archive_path)?;
let gz = flate2::read::GzDecoder::new(file);
let mut archive = tar::Archive::new(gz);
#[cfg(target_os = "linux")]
let lib_name = "libonnxruntime.so";
#[cfg(target_os = "macos")]
let lib_name = "libonnxruntime.dylib";
let mut extracted_real_path: Option<std::path::PathBuf> = None;
for entry in archive.entries()? {
let mut entry = entry?;
let entry_type = entry.header().entry_type();
let path = entry.path()?;
let name = path.to_string_lossy();
if entry_type == tar::EntryType::Symlink || entry_type == tar::EntryType::Link {
if name.contains(lib_name) {
tracing::debug!("Skipping symlink entry: {}", name);
}
continue;
}
let file_name = path
.file_name()
.map(|f| f.to_string_lossy().to_string())
.unwrap_or_default();
let is_target = file_name == lib_name
|| (file_name.starts_with("libonnxruntime")
&& file_name.contains(lib_name.trim_start_matches("libonnxruntime")));
let is_versioned = file_name.starts_with("libonnxruntime")
&& (file_name.ends_with(".dylib")
|| file_name.ends_with(".so")
|| file_name.contains(".so."));
if is_target || is_versioned {
let dest_path = dest_dir.join(&file_name);
entry.unpack(&dest_path)?;
tracing::info!("Extracted {}", file_name);
extracted_real_path = Some(dest_path);
}
}
if let Some(real_path) = extracted_real_path {
let canonical_path = dest_dir.join(lib_name);
if real_path != canonical_path {
fs::copy(&real_path, &canonical_path)?;
tracing::info!(
"Copied {} -> {} (avoiding dangling symlink)",
real_path.display(),
canonical_path.display()
);
}
let dest_path = &canonical_path;
#[cfg(target_os = "macos")]
{
let _ = std::process::Command::new("xattr")
.args(["-d", "com.apple.quarantine"])
.arg(dest_path)
.output();
match std::process::Command::new("codesign")
.args(["--force", "--deep", "-s", "-"])
.arg(dest_path)
.output()
{
Ok(output) if output.status.success() => {
tracing::info!("Ad-hoc signed {} for macOS Gatekeeper", lib_name);
}
Ok(output) => {
tracing::warn!(
"codesign returned non-zero for {}: {}",
lib_name,
String::from_utf8_lossy(&output.stderr)
);
}
Err(e) => {
tracing::warn!("codesign not available ({}), dlopen may fail on macOS", e);
}
}
}
return Ok(());
}
anyhow::bail!("{} not found in archive", lib_name);
}
}
pub fn ensure_downloaded(progress: Option<ProgressCallback>) -> Result<(PathBuf, PathBuf)> {
if let Ok(existing_path) = std::env::var("ORT_DYLIB_PATH") {
let path = PathBuf::from(&existing_path);
if path.exists() {
tracing::info!(
"Using existing ONNX Runtime from ORT_DYLIB_PATH: {:?}",
path
);
let models_dir = download_models(progress)?;
return Ok((models_dir, path));
}
}
let models_dir = download_models(progress.clone())?;
let onnx_path = download_onnx_runtime(progress)?;
std::env::set_var("ORT_DYLIB_PATH", &onnx_path);
tracing::info!("Set ORT_DYLIB_PATH to {:?}", onnx_path);
Ok((models_dir, onnx_path))
}
pub fn print_status() {
let cache_dir = get_cache_dir();
let models_downloaded = are_models_downloaded();
let ner_models_downloaded = are_ner_models_downloaded();
let onnx_downloaded = is_onnx_runtime_downloaded();
println!("Shodh-Memory Cache Status:");
println!(" Cache directory: {cache_dir:?}");
println!(" Embedding models downloaded: {models_downloaded}");
println!(" NER models downloaded: {ner_models_downloaded}");
println!(" ONNX Runtime downloaded: {onnx_downloaded}");
if models_downloaded {
let models_dir = get_models_dir();
println!(" Embedding model path: {models_dir:?}");
}
if ner_models_downloaded {
let ner_dir = get_ner_models_dir();
println!(" NER model path: {ner_dir:?}");
}
if onnx_downloaded {
if let Some(path) = get_onnx_runtime_path() {
println!(" ONNX Runtime path: {path:?}");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cache_dir() {
let cache_dir = get_cache_dir();
assert!(cache_dir.to_string_lossy().contains("shodh-memory"));
}
#[test]
fn test_models_dir() {
let models_dir = get_models_dir();
assert!(models_dir.to_string_lossy().contains("minilm-l6"));
}
}