use crate::error::GigasttError;
use crate::sha256::{Sha256, hex_lower};
#[cfg(feature = "net")]
use anyhow::Context;
use anyhow::Result;
use std::path::Path;
#[cfg(unix)]
#[cfg(feature = "net")]
use std::os::fd::AsRawFd;
use super::variant::ModelVariant;
#[cfg(feature = "net")]
use super::variant::PREQUANT_RELEASE_BASE;
pub(crate) fn hash_file_sha256(path: &Path) -> std::io::Result<String> {
use std::io::Read;
let mut file = std::fs::File::open(path)?;
let mut hasher = Sha256::new();
let mut buf = [0u8; 64 * 1024];
loop {
let n = file.read(&mut buf)?;
if n == 0 {
break;
}
hasher.update(&buf[..n]);
}
Ok(hex_lower(&hasher.finalize()))
}
pub(crate) fn verify_pinned_checksum(path: &Path, expected: &str) -> Result<(), GigasttError> {
let actual = hash_file_sha256(path).map_err(|e| GigasttError::ModelLoad {
path: path.display().to_string(),
source: Some(e.into()),
})?;
if actual != expected {
return Err(GigasttError::ModelLoad {
path: path.display().to_string(),
source: Some(format!("SHA-256 mismatch: expected {expected}, got {actual}").into()),
});
}
Ok(())
}
pub(super) fn home_dir() -> Option<std::path::PathBuf> {
#[cfg(unix)]
{
std::env::var_os("HOME").map(std::path::PathBuf::from)
}
#[cfg(windows)]
{
std::env::var_os("USERPROFILE").map(std::path::PathBuf::from)
}
}
pub fn default_model_dir() -> String {
home_dir()
.map(|h| {
h.join(".gigastt")
.join("models")
.to_string_lossy()
.into_owned()
})
.unwrap_or_else(|| ".gigastt/models".into())
}
pub fn default_punct_model_dir() -> String {
home_dir()
.map(|h| {
h.join(".gigastt")
.join("models")
.join("punct")
.to_string_lossy()
.into_owned()
})
.unwrap_or_else(|| ".gigastt/models/punct".into())
}
pub fn default_vad_model_dir() -> String {
home_dir()
.map(|h| {
h.join(".gigastt")
.join("models")
.join("vad")
.to_string_lossy()
.into_owned()
})
.unwrap_or_else(|| ".gigastt/models/vad".into())
}
#[cfg(unix)]
#[cfg(feature = "net")]
pub(super) fn acquire_download_lock(dir: &Path) -> Result<std::fs::File> {
let lock_path = dir.join(".download.lock");
let file = std::fs::OpenOptions::new()
.write(true)
.create(true)
.truncate(true)
.open(&lock_path)
.context("Failed to create download lock file")?;
let fd = file.as_raw_fd();
let ret = unsafe { libc::flock(fd, libc::LOCK_EX) };
if ret != 0 {
anyhow::bail!("Failed to acquire download lock (another process is downloading)");
}
Ok(file)
}
#[derive(Debug, PartialEq, Eq)]
pub enum VariantAction {
Use(ModelVariant),
Download(ModelVariant),
}
pub fn resolve_variant(
requested: Option<ModelVariant>,
existing: Option<ModelVariant>,
) -> VariantAction {
match (requested, existing) {
(Some(req), Some(ex)) if req == ex => VariantAction::Use(req),
(Some(req), _) => VariantAction::Download(req),
(None, Some(ex)) => VariantAction::Use(ex),
(None, None) => VariantAction::Download(ModelVariant::default()),
}
}
#[cfg(feature = "net")]
pub async fn ensure_model(model_dir: &str) -> Result<()> {
ensure_model_variant(None, model_dir).await?;
Ok(())
}
#[cfg(feature = "net")]
pub async fn ensure_model_variant(
requested: Option<ModelVariant>,
model_dir: &str,
) -> Result<ModelVariant> {
let dir = Path::new(model_dir);
let existing = ModelVariant::detect_in_dir(dir).filter(|&v| is_usable_present(v, dir));
let variant = match resolve_variant(requested, existing) {
VariantAction::Use(v) => {
tracing::info!("Using existing {v:?} model at {model_dir}");
return Ok(v);
}
VariantAction::Download(v) => v,
};
if let Some(other) = existing
&& other != variant
{
tracing::warn!(
"Model directory {model_dir} holds {other:?} files but {variant:?} was \
requested; downloading the {variant:?} set (variants are never mixed)"
);
}
std::fs::create_dir_all(dir).context("Failed to create model directory")?;
#[cfg(unix)]
let _lock = acquire_download_lock(dir)?;
if is_usable_present(variant, dir) {
tracing::info!("Model ({variant:?}) found at {model_dir} after lock acquisition");
return Ok(variant);
}
if variant.is_ctc() {
tracing::info!("Model ({variant:?}) not found, downloading from HuggingFace...");
for file in variant.download_files() {
download_file(variant, file, dir).await?;
}
} else {
tracing::info!(
"Model ({variant:?}) not found, downloading pre-quantized INT8 bundle from {PREQUANT_RELEASE_BASE}..."
);
for file in variant.prequantized_files() {
let final_dest = dir.join(file);
if final_dest.exists() {
continue;
}
let url = format!("{PREQUANT_RELEASE_BASE}/{file}");
let expected = variant.prequantized_checksum(file);
stream_to_partial_then_finalize(&url, &final_dest, expected, file).await?;
}
}
tracing::info!("Model download complete");
Ok(variant)
}
#[cfg(feature = "net")]
pub async fn ensure_fp32_model_variant(
requested: Option<ModelVariant>,
model_dir: &str,
) -> Result<ModelVariant> {
let dir = Path::new(model_dir);
let existing = ModelVariant::detect_in_dir(dir).filter(|&v| is_model_present(v, dir));
let variant = match resolve_variant(requested, existing) {
VariantAction::Use(v) => {
tracing::info!("Using existing FP32 {v:?} model at {model_dir}");
return Ok(v);
}
VariantAction::Download(v) => v,
};
std::fs::create_dir_all(dir).context("Failed to create model directory")?;
#[cfg(unix)]
let _lock = acquire_download_lock(dir)?;
if is_model_present(variant, dir) {
tracing::info!("FP32 model ({variant:?}) found at {model_dir} after lock");
return Ok(variant);
}
tracing::info!("Downloading FP32 {variant:?} model set from HuggingFace...");
for file in variant.download_files() {
download_file(variant, file, dir).await?;
}
tracing::info!("FP32 model download complete");
Ok(variant)
}
#[cfg(feature = "net")]
pub async fn ensure_prequantized_model_variant(
requested: Option<ModelVariant>,
model_dir: &str,
) -> Result<ModelVariant> {
let dir = Path::new(model_dir);
let variant = requested
.or_else(|| ModelVariant::detect_in_dir(dir).filter(|&v| is_usable_present(v, dir)))
.unwrap_or_default();
if is_prequantized_present(variant, dir) {
tracing::info!("Using existing {variant:?} INT8 model at {model_dir}");
return Ok(variant);
}
std::fs::create_dir_all(dir).context("Failed to create model directory")?;
#[cfg(unix)]
let _lock = acquire_download_lock(dir)?;
if is_prequantized_present(variant, dir) {
tracing::info!("Pre-quantized {variant:?} model found at {model_dir} after lock");
return Ok(variant);
}
tracing::info!("Downloading pre-quantized {variant:?} model from {PREQUANT_RELEASE_BASE}...");
for file in variant.prequantized_files() {
let final_dest = dir.join(file);
if final_dest.exists() {
continue;
}
let url = format!("{PREQUANT_RELEASE_BASE}/{file}");
let expected = variant.prequantized_checksum(file);
stream_to_partial_then_finalize(&url, &final_dest, expected, file).await?;
}
tracing::info!("Pre-quantized model download complete");
Ok(variant)
}
pub fn is_model_present(variant: ModelVariant, dir: &Path) -> bool {
variant
.download_files()
.iter()
.all(|f| dir.join(f).exists())
}
pub fn is_prequantized_present(variant: ModelVariant, dir: &Path) -> bool {
variant
.prequantized_files()
.iter()
.all(|f| dir.join(f).exists())
}
pub fn is_usable_present(variant: ModelVariant, dir: &Path) -> bool {
if variant.is_ctc() {
return is_model_present(variant, dir) || is_prequantized_present(variant, dir);
}
is_prequantized_present(variant, dir)
}
#[cfg(test)]
pub(super) fn partial_path(final_path: &Path) -> std::path::PathBuf {
let mut s: std::ffi::OsString = final_path.as_os_str().to_owned();
s.push(".partial");
std::path::PathBuf::from(s)
}
#[cfg(feature = "net")]
pub(super) mod fetch;
#[cfg(feature = "net")]
use fetch::{download_file, stream_to_partial_then_finalize};
#[cfg(feature = "ane")]
pub(super) mod ane;
#[cfg(all(feature = "net", feature = "ane"))]
pub use ane::ensure_ane_packages;
#[cfg(feature = "ane")]
pub use ane::{ane_package_complete, ane_package_dir_name, default_ane_model_dir, is_ane_present};
#[cfg(feature = "net")]
mod sidecars;
#[cfg(all(feature = "net", feature = "diarization"))]
pub use sidecars::ensure_speaker_model;
#[cfg(feature = "net")]
pub use sidecars::{ensure_punct_model, ensure_vad_model};