use std::collections::HashMap;
use std::path::Path;
use std::sync::{Arc, Mutex};
use crate::model::{ANE_BUCKETS, ane_package_complete, ane_package_dir_name};
use crate::runtime::{error::RuntimeError, factory::Runtime, session::RuntimeSession};
use super::bridge;
use super::encoder_session::{AneEncoderSession, BucketModel, SharedModel};
pub struct AneRuntime {
ort: Box<dyn Runtime>,
bucket_cache: Arc<Mutex<HashMap<usize, Arc<SharedModel>>>>,
}
impl AneRuntime {
pub fn new(ort: Box<dyn Runtime>) -> Self {
Self {
ort,
bucket_cache: Arc::new(Mutex::new(HashMap::new())),
}
}
fn shared_bucket_model(
&self,
ane_dir: &Path,
bucket: usize,
) -> Result<Arc<SharedModel>, RuntimeError> {
let mut cache = self
.bucket_cache
.lock()
.map_err(|_| RuntimeError::InferenceFailed("ANE bucket cache poisoned".into()))?;
if let Some(model) = cache.get(&bucket) {
return Ok(Arc::clone(model));
}
let pkg = ane_dir.join(ane_package_dir_name(bucket));
tracing::info!(bucket, package = %pkg.display(), "compiling ANE encoder bucket (cold-start, may take several seconds)");
let model = bridge::compile_and_load(&pkg, true)?;
let shared = Arc::new(SharedModel(model));
cache.insert(bucket, Arc::clone(&shared));
Ok(shared)
}
}
impl Runtime for AneRuntime {
fn load_session(
&self,
model_path: &Path,
is_encoder: bool,
) -> Result<Box<dyn RuntimeSession>, RuntimeError> {
if !is_encoder {
return self.ort.load_session(model_path, false);
}
let ane_dir =
model_path
.parent()
.map(|p| p.join("ane"))
.ok_or_else(|| RuntimeError::LoadFailed {
path: model_path.to_path_buf(),
message: "encoder model path has no parent directory".to_string(),
})?;
let available: Vec<usize> = ANE_BUCKETS
.iter()
.copied()
.filter(|&b| ane_package_complete(&ane_dir.join(ane_package_dir_name(b))))
.collect();
if available.is_empty() {
return Err(RuntimeError::LoadFailed {
path: model_path.to_path_buf(),
message: format!(
"ANE encoder packages not found in {}; run `gigastt download --ane` (or convert locally)",
ane_dir.display()
),
});
}
let mut buckets = Vec::with_capacity(available.len());
for bucket in available {
let model = self.shared_bucket_model(&ane_dir, bucket)?;
buckets.push(BucketModel {
size: bucket,
model,
});
}
let ort_fallback = self.ort.load_session(model_path, true)?;
Ok(Box::new(AneEncoderSession::new(buckets, ort_fallback)))
}
}