use anyhow::Result;
use std::path::Path;
use super::super::variant::ANE_BUCKETS;
#[cfg(feature = "net")]
use super::super::variant::{ANE_RELEASE_BASE, ANE_TAR_CHECKSUMS};
use super::home_dir;
#[cfg(all(unix, feature = "net"))]
use super::acquire_download_lock;
#[cfg(feature = "net")]
use super::fetch::stream_to_partial_then_finalize;
#[cfg(feature = "net")]
use anyhow::Context;
#[cfg(feature = "ane")]
pub fn ane_package_dir_name(bucket: usize) -> String {
format!("gigaam_v3_encoder_{bucket}.mlpackage")
}
#[cfg(all(feature = "net", feature = "ane"))]
pub(crate) fn ane_tar_name(bucket: usize) -> String {
format!("{}.tar", ane_package_dir_name(bucket))
}
#[cfg(all(feature = "net", feature = "ane"))]
fn ane_tar_checksum(bucket: usize) -> Option<&'static str> {
ANE_TAR_CHECKSUMS
.iter()
.find(|(b, _)| *b == bucket)
.and_then(|(_, sum)| if sum.is_empty() { None } else { Some(*sum) })
}
#[cfg(feature = "ane")]
pub fn default_ane_model_dir() -> String {
home_dir()
.map(|h| {
h.join(".gigastt")
.join("models")
.join("ane")
.to_string_lossy()
.into_owned()
})
.unwrap_or_else(|| ".gigastt/models/ane".into())
}
#[cfg(feature = "ane")]
pub fn ane_package_complete(pkg_dir: &Path) -> bool {
pkg_dir.is_dir()
&& pkg_dir.join("Manifest.json").is_file()
&& pkg_dir
.join("Data")
.join("com.apple.CoreML")
.join("model.mlmodel")
.is_file()
&& pkg_dir
.join("Data")
.join("com.apple.CoreML")
.join("weights")
.join("weight.bin")
.is_file()
}
#[cfg(feature = "ane")]
pub fn is_ane_present(dir: &Path) -> bool {
ANE_BUCKETS
.iter()
.all(|&b| ane_package_complete(&dir.join(ane_package_dir_name(b))))
}
#[cfg(all(feature = "net", feature = "ane"))]
pub async fn ensure_ane_packages(model_dir: &str) -> Result<()> {
let dir = Path::new(model_dir);
if is_ane_present(dir) {
tracing::info!("ANE encoder packages found at {model_dir}");
return Ok(());
}
std::fs::create_dir_all(dir).context("Failed to create ANE model directory")?;
#[cfg(unix)]
let _lock = acquire_download_lock(dir)?;
if is_ane_present(dir) {
tracing::info!("ANE encoder packages found at {model_dir} after lock");
return Ok(());
}
tracing::info!("Downloading ANE encoder packages from {ANE_RELEASE_BASE}...");
for &bucket in ANE_BUCKETS {
let pkg_name = ane_package_dir_name(bucket);
if ane_package_complete(&dir.join(&pkg_name)) {
continue;
}
let checksum = require_ane_tar_checksum(bucket)?;
let tar_name = ane_tar_name(bucket);
let tar_dest = dir.join(&tar_name);
let url = format!("{ANE_RELEASE_BASE}/{tar_name}");
stream_to_partial_then_finalize(&url, &tar_dest, Some(checksum), &tar_name).await?;
tracing::info!("Unpacking {tar_name} into {model_dir}");
if let Err(e) = extract_ane_tar_atomic(&tar_dest, dir, &pkg_name) {
let _ = std::fs::remove_file(&tar_dest);
return Err(e);
}
std::fs::remove_file(&tar_dest)
.with_context(|| format!("Failed to remove {}", tar_dest.display()))?;
}
tracing::info!("ANE encoder packages download complete");
Ok(())
}
#[cfg(all(feature = "net", feature = "ane"))]
pub(crate) fn extract_ane_tar_atomic(tar_dest: &Path, dir: &Path, pkg_name: &str) -> Result<()> {
let stamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos();
let staging = dir.join(format!(".extract.{}.{}", std::process::id(), stamp));
let cleanup_staging = || {
let _ = std::fs::remove_dir_all(&staging);
};
if let Err(e) = std::fs::create_dir_all(&staging)
.with_context(|| format!("Failed to create staging dir {}", staging.display()))
{
cleanup_staging();
return Err(e);
}
let unpack = (|| -> Result<()> {
let tar_file = std::fs::File::open(tar_dest)
.with_context(|| format!("Failed to open {}", tar_dest.display()))?;
tar::Archive::new(tar_file)
.unpack(&staging)
.with_context(|| format!("Failed to unpack {}", tar_dest.display()))?;
let src = staging.join(pkg_name);
let dest = dir.join(pkg_name);
if dest.exists() {
std::fs::remove_dir_all(&dest)
.with_context(|| format!("Failed to remove stale {}", dest.display()))?;
}
std::fs::rename(&src, &dest)
.with_context(|| format!("Failed to rename {} -> {}", src.display(), dest.display()))?;
Ok(())
})();
cleanup_staging();
unpack
}
#[cfg(all(feature = "net", feature = "ane"))]
pub(crate) fn require_ane_tar_checksum(bucket: usize) -> Result<&'static str> {
ane_tar_checksum(bucket).ok_or_else(|| {
anyhow::anyhow!(
"ANE encoder release not yet published; run the Release ANE workflow \
(release-ane.yml), then pin the per-bucket .tar SHA-256 from \
SHA256SUMS.txt in ANE_TAR_CHECKSUMS"
)
})
}