use sha2::{Digest, Sha256};
use std::fs::{File, OpenOptions};
use std::io::{Read, Write};
use std::path::{Path, PathBuf};
#[derive(Clone, Copy, Debug)]
pub struct ModelInfo {
pub name: &'static str,
pub backend: &'static str,
pub filename: &'static str,
pub url: &'static str,
pub revision: &'static str,
pub sha256: &'static str,
pub license: &'static str,
pub sample_rate: u32,
}
pub const MODELS: &[ModelInfo] = &[ModelInfo {
name: "gtcrn-dns3",
backend: "gtcrn",
filename: "gtcrn_simple.onnx",
url: "https://raw.githubusercontent.com/Xiaobin-Rong/gtcrn/3862c44808dca492ea5a8a145d2dc2a1028d08c8/stream/onnx_models/gtcrn_simple.onnx",
revision: "3862c44808dca492ea5a8a145d2dc2a1028d08c8",
sha256: "b4718df6228e7bdf1a8a435cf98f838636eb2fd331acabf86ba87c5192ebcb87",
license: "MIT",
sample_rate: 16_000,
}];
pub fn find(name: &str) -> Option<&'static ModelInfo> {
MODELS
.iter()
.find(|model| model.name == name || model.backend == name)
}
pub fn cache_dir() -> Result<PathBuf, String> {
if let Some(path) = std::env::var_os("DENOIZE_MODEL_DIR") {
return Ok(PathBuf::from(path));
}
#[cfg(target_os = "windows")]
if let Some(path) = std::env::var_os("LOCALAPPDATA") {
return Ok(PathBuf::from(path).join("denoize").join("models"));
}
if let Some(path) = std::env::var_os("XDG_CACHE_HOME") {
return Ok(PathBuf::from(path).join("denoize").join("models"));
}
std::env::var_os("HOME")
.map(|path| PathBuf::from(path).join(".cache/denoize/models"))
.ok_or_else(|| "cannot locate model cache; set DENOIZE_MODEL_DIR".into())
}
pub fn path(model: &ModelInfo) -> Result<PathBuf, String> {
Ok(cache_dir()?.join(model.name).join(model.filename))
}
pub fn verify(model: &ModelInfo) -> Result<PathBuf, String> {
let path = path(model)?;
if !path.is_file() {
return Err(format!("model is not installed: {}", path.display()));
}
let actual = sha256(&path)?;
if actual != model.sha256 {
return Err(format!(
"checksum mismatch for {}: expected {}, got {}",
path.display(),
model.sha256,
actual
));
}
Ok(path)
}
pub fn install(model: &ModelInfo) -> Result<PathBuf, String> {
install_with_progress(model, || false, |_, _| {})
}
pub fn install_with_progress<C, P>(
model: &ModelInfo,
mut cancelled: C,
mut progress: P,
) -> Result<PathBuf, String>
where
C: FnMut() -> bool,
P: FnMut(u64, Option<u64>),
{
if let Ok(path) = verify(model) {
return Ok(path);
}
let destination = path(model)?;
let parent = destination
.parent()
.ok_or_else(|| "invalid model cache path".to_string())?;
std::fs::create_dir_all(parent)
.map_err(|error| format!("failed to create {}: {error}", parent.display()))?;
let partial = destination.with_extension("onnx.part");
let downloaded = partial.metadata().map(|meta| meta.len()).unwrap_or(0);
let mut request = ureq::get(model.url).set("User-Agent", "denoize-model-manager");
if downloaded > 0 {
request = request.set("Range", &format!("bytes={downloaded}-"));
}
let response = request
.call()
.map_err(|error| format!("failed to download {}: {error}", model.url))?;
let resumed = downloaded > 0 && response.status() == 206;
let response_length = response
.header("Content-Length")
.and_then(|value| value.parse::<u64>().ok());
let total = response_length.map(|length| if resumed { downloaded + length } else { length });
let mut received = if resumed { downloaded } else { 0 };
progress(received, total);
let mut output = OpenOptions::new()
.create(true)
.write(true)
.append(resumed)
.truncate(!resumed)
.open(&partial)
.map_err(|error| format!("failed to open {}: {error}", partial.display()))?;
let mut reader = response.into_reader();
let mut buffer = [0_u8; 64 * 1024];
loop {
if cancelled() {
output
.flush()
.map_err(|error| format!("failed to flush {}: {error}", partial.display()))?;
return Err("cancelled".into());
}
let count = reader
.read(&mut buffer)
.map_err(|error| format!("failed to download {}: {error}", model.url))?;
if count == 0 {
break;
}
output
.write_all(&buffer[..count])
.map_err(|error| format!("failed to save {}: {error}", partial.display()))?;
received += count as u64;
progress(received, total);
}
output
.flush()
.map_err(|error| format!("failed to flush {}: {error}", partial.display()))?;
let actual = sha256(&partial)?;
if actual != model.sha256 {
return Err(format!(
"downloaded model checksum mismatch: expected {}, got {} (partial kept at {})",
model.sha256,
actual,
partial.display()
));
}
std::fs::rename(&partial, &destination).map_err(|error| {
format!(
"failed to move {} to {}: {error}",
partial.display(),
destination.display()
)
})?;
Ok(destination)
}
pub fn remove(model: &ModelInfo) -> Result<bool, String> {
let destination = path(model)?;
let partial = destination.with_extension("onnx.part");
let removed = remove_file_if_present(&destination)? | remove_file_if_present(&partial)?;
if let Some(directory) = destination.parent() {
match std::fs::remove_dir(directory) {
Ok(()) => {}
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {}
Err(error) if error.kind() == std::io::ErrorKind::DirectoryNotEmpty => {}
Err(error) => return Err(format!("failed to remove {}: {error}", directory.display())),
}
}
Ok(removed)
}
pub fn update(model: &ModelInfo) -> Result<PathBuf, String> {
update_with_progress(model, || false, |_, _| {})
}
pub fn update_with_progress<C, P>(
model: &ModelInfo,
cancelled: C,
progress: P,
) -> Result<PathBuf, String>
where
C: FnMut() -> bool,
P: FnMut(u64, Option<u64>),
{
let destination = path(model)?;
let backup = destination.with_extension("onnx.backup");
let had_existing = destination.is_file();
if had_existing {
std::fs::rename(&destination, &backup).map_err(|error| {
format!(
"failed to stage existing model {}: {error}",
destination.display()
)
})?;
}
match install_with_progress(model, cancelled, progress) {
Ok(path) => {
let _ = std::fs::remove_file(backup);
Ok(path)
}
Err(error) => {
if had_existing {
let _ = std::fs::remove_file(&destination);
std::fs::rename(&backup, &destination).map_err(|restore_error| {
format!("{error}; additionally failed to restore old model: {restore_error}")
})?;
}
Err(error)
}
}
}
fn remove_file_if_present(path: &Path) -> Result<bool, String> {
match std::fs::remove_file(path) {
Ok(()) => Ok(true),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(false),
Err(error) => Err(format!("failed to remove {}: {error}", path.display())),
}
}
fn sha256(path: &Path) -> Result<String, String> {
let mut input =
File::open(path).map_err(|error| format!("failed to open {}: {error}", path.display()))?;
let mut digest = Sha256::new();
let mut buffer = [0_u8; 64 * 1024];
loop {
let count = input
.read(&mut buffer)
.map_err(|error| format!("failed to read {}: {error}", path.display()))?;
if count == 0 {
break;
}
digest.update(&buffer[..count]);
}
Ok(format!("{:x}", digest.finalize()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn manifest_has_pinned_integrity_and_metadata() {
for model in MODELS {
assert_eq!(model.sha256.len(), 64);
assert_eq!(model.revision.len(), 40);
assert!(model.url.contains(model.revision));
assert!(model.sample_rate > 0);
assert!(!model.license.is_empty());
}
}
#[test]
fn removal_is_idempotent() {
let directory =
std::env::temp_dir().join(format!("denoize-model-remove-test-{}", std::process::id()));
std::fs::create_dir_all(&directory).unwrap();
let path = directory.join("model.onnx");
std::fs::write(&path, b"model").unwrap();
assert!(remove_file_if_present(&path).unwrap());
assert!(!remove_file_if_present(&path).unwrap());
std::fs::remove_dir(directory).unwrap();
}
}